melghorab/code-assistant
0
1import os2 3os.environ["CHAINLIT_DISABLE_WEBSOCKETS"] = "true"4 # Also consider setting these for HF Spaces5os.environ["CHAINLIT_SERVER_PORT"] = "7860"6os.environ["CHAINLIT_SERVER_HOST"] = "0.0.0.0"7os.environ["CHAINLIT_USE_PREDEFINED_HOST_PORT"] = "true"8os.environ["CHAINLIT_USE_HTTP"] = "true"9import getpass10from operator import itemgetter11from typing import List, Dict12import json13import requests14 15 16 17#LangChain, LangGraph18from langchain_openai import ChatOpenAI19from langgraph.graph import START, StateGraph, END20from typing_extensions import List, TypedDict21from langchain_core.documents import Document22from langchain_core.prompts import ChatPromptTemplate23from langchain.schema.output_parser import StrOutputParser24from langchain_core.tools import Tool, tool25from langgraph.prebuilt import ToolNode26from typing import TypedDict, Annotated27from langgraph.graph.message import add_messages28import operator29from langchain_core.messages import BaseMessage, HumanMessage, AIMessage30from langchain.vectorstores import Qdrant31from langchain.embeddings import OpenAIEmbeddings32from langchain.schema import Document33from qdrant_client import QdrantClient34from qdrant_client.http.models import Distance, VectorParams35 36import chainlit as cl37import tempfile38import shutil39 40 41 42 43#helper imports44from code_analysis import *45from tools import search_pypi, write_to_docx46from prompts import describe_imports, main_prompt, documenter_prompt47from states import AgentState48 49if os.environ.get("SPACE_ID"): # Check if running on HF Spaces50 os.environ["CHAINLIT_DISABLE_WEBSOCKETS"] = "true"51 # Also consider setting these for HF Spaces52 os.environ["CHAINLIT_SERVER_PORT"] = "7860"53 os.environ["CHAINLIT_SERVER_HOST"] = "0.0.0.0"54 55 56# Global variables to store processed data57processed_file_path = None58document_file_path = None59vectorstore = None60main_chain = None61qdrant_client = None62 63@cl.on_chat_start64async def on_chat_start():65 print("Chat session started")66 67 await cl.Message(content="Welcome to the Python Code Documentation Assistant! Please upload a Python file to get started.").send()68 69@cl.on_message70async def on_message(message: cl.Message):71 global processed_file_path, document_file_path, vectorstore, main_chain, qdrant_client72 73 if message.elements and any(el.type == "file" for el in message.elements):74 file_elements = [el for el in message.elements if el.type == "file"]75 file_element = file_elements[0]76 is_python_file = (77 file_element.mime.startswith("text/x-python") or 78 file_element.name.endswith(".py") or79 file_element.mime == "text/plain" # Some systems identify .py as text/plain80 )81 if is_python_file:82 # Send processing message83 msg = cl.Message(content="Processing your Python file...")84 await msg.send()85 86 print(f'file element \n {file_element} \n')87 88 # Save uploaded file to a temporary location89 temp_dir = tempfile.mkdtemp()90 file_path = os.path.join(temp_dir, file_element.name)91 92 with open(file_element.path, "rb") as source_file:93 file_content_bytes = source_file.read()94 with open(file_path, "wb") as destination_file:95 destination_file.write(file_content_bytes)96 97 processed_file_path = file_path98 99 try:100 101 # read file and extract imports102 file_content = read_python_file(file_path)103 imports = extract_imports(file_content, file_path)104 105 print(f'Done reading file')106 107 # Define describe packages graph108 search_packages_tools = [search_pypi]109 describe_imports_llm = ChatOpenAI(model="gpt-4o-mini")110 # describe_imports_llm = describe_imports_llm.bind_tools(tools = search_packages_tools, tool_choice="required")111 112 describe_imports_prompt = ChatPromptTemplate.from_messages([113 ("system", describe_imports),114 ("human", "{imports}")115 ])116 117 describe_imports_chain = (118 {"code_language": itemgetter("code_language"), "imports": itemgetter("imports")}119 | describe_imports_prompt | describe_imports_llm | StrOutputParser()120 )121 122 print(f'done defining imports chain')123 124 125 # Define imports chain function126 def call_imports_chain(state):127 last_message= state["messages"][-1]128 content = json.loads(last_message.content)129 chain_input = {"code_language": content['code_language'], 130 "imports": content['imports']}131 response = describe_imports_chain.invoke(chain_input)132 return {"messages": [AIMessage(content=response)]}133 134 # bind model to tool or ToolNode135 imports_tool_node = ToolNode(search_packages_tools)136 137 # construct graph and compile138 uncompiled_imports_graph = StateGraph(AgentState)139 uncompiled_imports_graph.add_node("imports_agent", call_imports_chain)140 uncompiled_imports_graph.add_node("imports_action", imports_tool_node)141 uncompiled_imports_graph.set_entry_point("imports_agent")142 143 def should_continue(state):144 last_message = state["messages"][-1]145 146 if last_message.tool_calls:147 return "imports_action"148 149 return END150 151 uncompiled_imports_graph.add_conditional_edges(152 "imports_agent",153 should_continue154 )155 156 uncompiled_imports_graph.add_edge("imports_action", "imports_agent")157 158 compiled_imports_graph = uncompiled_imports_graph.compile()159 160 print(f'compiled imports graph')161 # Invoke imports graph162 initial_state = {163 "messages": [{164 "role": "human",165 "content": json.dumps({166 "code_language": "python",167 "imports": imports168 })169 }]170 }171 172 # await msg.update(content="Analyzing imports and generating documentation...")173 msg.content = "Analyzing your code and generating documentation..."174 await msg.update()175 176 msg = cl.Message(content="Analyzing your code and generating documentation...")177 await msg.send()178 179 result = compiled_imports_graph.invoke(initial_state)180 181 # Define qdrant Database182 qdrant_client = QdrantClient(":memory:")183 184 embedding_model = OpenAIEmbeddings(model="text-embedding-3-small")185 embedding_dim = 1536186 187 qdrant_client.create_collection(188 collection_name="description_rag_data",189 vectors_config=VectorParams(size=embedding_dim, distance=Distance.COSINE),190 )191 192 vectorstore = Qdrant(qdrant_client, collection_name="description_rag_data", embeddings=embedding_model)193 194 # Add packages chunks195 text = result['messages'][-1].content196 chunks = [197 {"type": "Imported Packages", "name": "Imported Packages", "content": text},198 #{"type": "Source Code", "name": "Source Code", "content": file_content},199 200 ]201 202 docs = [203 Document(204 page_content=f"{chunk['type']} - {chunk['name']} - {chunk['content']}", # Content for the model205 metadata={**chunk} # Store metadata, but don't put embeddings here206 )207 for chunk in chunks208 ]209 vectorstore.add_documents(docs)210 qdrant_retriever = vectorstore.as_retriever(search_kwargs={"k": 3})211 212 print('done adding docs to DB')213 #define documenter chain214 documenter_llm = ChatOpenAI(model="gpt-4o-mini")215 documenter_llm_prompt = ChatPromptTemplate.from_messages([216 ("system", documenter_prompt),217 ])218 documenter_chain = (219 {"context": itemgetter("context")}220 | documenter_llm_prompt221 | documenter_llm222 | StrOutputParser()223 )224 225 print('done defining documenter chain')226 #extract description chunks from database227 collection_name = "description_rag_data"228 all_points = qdrant_client.scroll(collection_name=collection_name, limit=1000)[0] # Adjust limit if needed229 one_chunk = all_points[0].payload230 input_text = f"type: {one_chunk['metadata']['type']} \nname: {one_chunk['metadata']['name']} \ncontent: {one_chunk['metadata']['content']}"231 232 print('done extracting chunks form DB')233 234 document_response = documenter_chain.invoke({"context": input_text})235 236 print('done invoking documenter chain and will write in docx')237 # write packages description in word file238 document_file_path = write_to_docx(document_response)239 240 print('done writing docx file')241 # Set up Main Chain for chat242 main_llm = ChatOpenAI(model="gpt-4o-mini", temperature=0)243 244 245 main_llm_prompt = ChatPromptTemplate.from_messages([246 ("system", main_prompt),247 ("human", "{query}")248 ])249 250 main_chain = (251 {"context": itemgetter("query") | qdrant_retriever, "code_language": itemgetter("code_language"), "query": itemgetter("query"), }252 | main_llm_prompt253 | main_llm254 | StrOutputParser()255 )256 257 print('done defining main chain')258 # Present download button for the document259 elements = [260 cl.File(261 name="documentation.docx",262 path=document_file_path,263 display="inline"264 )265 ]266 print('done defining elements')267 msg.content = "โ
Your Python file has been processed! You can download the documentation file below. How can I help you with your code?"268 msg.elements = elements269 await msg.update()270 271 except Exception as e:272 msg.content = f"โ Error processing file: {str(e)}"273 await msg.update()274 275 276 else:277 await cl.Message(content="Please upload a Python (.py) file.").send()278 279 # Handle chat messages if file has been processed280 elif processed_file_path and main_chain:281 user_input = message.content282 # Send thinking message283 msg = cl.Message(content="Thinking...")284 await msg.send()285 286 try:287 # Use main_chain to answer the query288 # invoke main chain289 inputs = {290 'code_language': 'Python',291 'query': user_input292 }293 294 response = main_chain.invoke(inputs)295 296 # Update with the response297 msg.content = response298 await msg.update()299 300 301 except Exception as e:302 msg.content = f"โ Error processing your question: {str(e)}"303 await msg.update()304 305 306 307 else:308 await cl.Message(content="Please upload a Python file first before asking questions.").send()309 310 311@cl.on_stop312def on_stop():313 global processed_file_path314 # Clean up temporary files315 if processed_file_path and os.path.exists(os.path.dirname(processed_file_path)):316 shutil.rmtree(os.path.dirname(processed_file_path))317 318 