Ransaka/Code-Assistant
0
1import os2from typing import List3import numpy as np4import redis5import google.generativeai as genai6from tqdm import tqdm7import time8 9from redis.commands.search.field import (10 TagField,11 TextField,12 VectorField,13)14from redis.commands.search.indexDefinition import IndexDefinition, IndexType15from redis.commands.search.query import Query16from sourcegraph import Sourcegraph17 18 19INDEX_NAME = "idx:codes_vss"20 21genai.configure(api_key=os.environ["GEMINI_API_KEY"])22 23generation_config = {24 "temperature": 1,25 "top_p": 0.95,26 "top_k": 64,27 "max_output_tokens": 8192,28 "response_mime_type": "text/plain",29}30 31model = genai.GenerativeModel(32 model_name="gemini-1.5-flash",33 generation_config=generation_config,34 system_instruction="You are optimized to generate accurate descriptions for given Python codes. When the user inputs the code, you must return the description according to its goal and functionality. You are not allowed to generate additional details. The user expects at least 5 sentence-long descriptions.",35)36 37def fetch_data(url):38 def get_description(code):39 chat_session = model.start_chat(40 history=[41 {42 "role": "user",43 "parts": [44 f"Code: {code}",45 ],46 },47 ]48 )49 response = chat_session.send_message("INSERT_INPUT_HERE")50 51 return response.text52 gihub_repository = Sourcegraph(url)53 gihub_repository.run()54 data = dict(gihub_repository.node_data)55 for key, value in tqdm(data.items()):56 data[key]['description'] = get_description(value['definition'])57 data[key]['uses'] = ", ".join(list(gihub_repository.get_dependencies(key)))58 time.sleep(3) #to overcome limit issues59 return data60 61def get_embeddings(content: List):62 return genai.embed_content(model='models/text-embedding-004',content=content)['embedding']63 64def ingest_data(client: redis.Redis, data):65 try:66 client.delete(client.keys("code:*"))67 except:68 pass69 pipeline = client.pipeline()70 for i, code_metadata in enumerate(data.values(), start=1):71 redis_key = f"code:{i:03}"72 pipeline.json().set(redis_key, "$", code_metadata)73 _ = pipeline.execute()74 keys = sorted(client.keys("code:*"))75 defs = client.json().mget(keys, "$.definition")76 descs = client.json().mget(keys, "$.description")77 embed_inputs = []78 79 for i in range(1, len(keys)+1):80 embed_inputs.append(81 f"""{defs[i-1][0]}\n\n{descs[i-1][0]}"""82 )83 embeddings = get_embeddings(embed_inputs)84 VECTOR_DIMENSION = len(embeddings[0])85 pipeline = client.pipeline()86 for key, embedding in zip(keys, embeddings):87 pipeline.json().set(key, "$.embeddings", embedding)88 pipeline.execute()89 90 schema = (91 TextField("$.name", no_stem=True, as_name="name"),92 TagField("$.type", as_name="type"),93 TextField("$.definition", no_stem=True, as_name="definition"),94 TextField("$.file_name", no_stem=True, as_name="file_name"),95 TextField("$.description", no_stem=True, as_name="description"),96 TextField("$.uses", no_stem=True, as_name="uses"),97 VectorField(98 "$.embeddings",99 "HNSW",100 {101 "TYPE": "FLOAT32",102 "DIM": VECTOR_DIMENSION,103 "DISTANCE_METRIC": "COSINE",104 },105 as_name="vector",106 ),107 )108 definition = IndexDefinition(prefix=["code:"], index_type=IndexType.JSON)109 try:110 _ = client.ft(INDEX_NAME).create_index(fields=schema, definition=definition)111 except redis.exceptions.ResponseError:112 client.ft(INDEX_NAME).dropindex()113 _ = client.ft(INDEX_NAME).create_index(fields=schema, definition=definition)114 info = client.ft(INDEX_NAME).info()115 num_docs = info["num_docs"]116 indexing_failures = info["hash_indexing_failures"]117 return f"{num_docs} documents indexed with {indexing_failures} failures"