Team Ai
Apppublic

melghorab/code-assistant

sourceHugging Faceopenrailupdated 2y agoView on Hugging Face
0likes
app_2.py392 linesDownload Raw Back to root
1import os2import getpass3from operator import itemgetter4from typing import List, Dict5import json6import requests7import traceback8 9 10 11#LangChain, LangGraph12from langchain_openai import ChatOpenAI13from langgraph.graph import START, StateGraph, END14from typing_extensions import List, TypedDict15# from langchain_core.documents import Document16from langchain_core.prompts import ChatPromptTemplate17from langchain.schema.output_parser import StrOutputParser18from langchain_core.tools import Tool, tool19from langgraph.prebuilt import ToolNode20from typing import TypedDict, Annotated21from langgraph.graph.message import add_messages22import operator23from langchain_core.messages import BaseMessage, HumanMessage, AIMessage, SystemMessage24from langchain.vectorstores import Qdrant25from langchain.embeddings import OpenAIEmbeddings26from langchain.schema import Document27from qdrant_client import QdrantClient28from qdrant_client.http.models import Distance, VectorParams29 30 31import chainlit as cl32import tempfile33import shutil34 35#helper imports36from code_analysis import *37from tools import search_pypi, write_to_docx38from prompts import main_prompt, documenter_prompt, code_description_prompt39from states import AgentState40 41 42 43# Global variables to store processed data44processed_file_path = None45document_file_path = None46vectorstore = None47main_chain = None48qdrant_client = None49 50@cl.on_chat_start51async def on_chat_start():52    await cl.Message(content="Welcome to the Python Code Documentation Assistant! Please upload a Python file to get started.").send()53 54@cl.on_message55async def on_message(message: cl.Message):56    global processed_file_path, document_file_path, vectorstore, main_chain, qdrant_client57    58    if message.elements and any(el.type == "file" for el in message.elements):59        file_elements = [el for el in message.elements if el.type == "file"]60        file_element = file_elements[0]61        is_python_file = (62            file_element.mime.startswith("text/x-python") or 63            file_element.name.endswith(".py") or64            file_element.mime == "text/plain"  # Some systems identify .py as text/plain65        )66        if is_python_file:67            # Send processing message68            msg = cl.Message(content="Processing your Python file...")69            await msg.send()70 71            print(f'file element \n {file_element} \n')72 73            # Save uploaded file to a temporary location74            temp_dir = tempfile.mkdtemp()75            file_path = os.path.join(temp_dir, file_element.name)76 77            with open(file_element.path, "rb") as source_file:78                file_content_bytes = source_file.read()79                with open(file_path, "wb") as destination_file:80                    destination_file.write(file_content_bytes)81                82            processed_file_path = file_path83 84            try:85                86                # read file and extract imports87                file_content = read_python_file(file_path)88                # imports = extract_imports(file_content, file_path)89 90                print(f'Done reading file')91 92                # Define describe packages graph93                search_packages_tools = [search_pypi]94##################################### DESCRIBE CODE AGENT ####################################95                describe_code_llm = ChatOpenAI(model="gpt-4o-mini")96                # describe_imports_llm = describe_imports_llm.bind_tools(tools = search_packages_tools, tool_choice="required")97 98                describe_code_prompt = ChatPromptTemplate.from_messages([99                        ("system", code_description_prompt),100                        ("human", "{code}")101                    ])102 103                describe_code_chain = (104                    {"code_language": itemgetter("code_language"), "code": itemgetter("code")}105                    | describe_code_prompt | describe_code_llm | StrOutputParser()106                )107 108                print(f'done defining imports chain')109 110 111                # Define describe code chain node112                def describe_code(state):113                    # print("Starting chain function")114                    last_message= state["messages"][-1]115                    # print(f'last message is \n {last_message}')116                    content = json.loads(last_message.content)117                    # print(f'content is {content}')118                    # print(type(content))119                    chain_input = {"code_language": content['code_language'], 120                                    "code": content['code']}121                    # print(f'chain_input is {chain_input}')122                    # print(type(chain_input))123                    response = describe_code_chain.invoke(chain_input)124                    # print(f"Chain response: {response}")125                    return {"messages": [AIMessage(content=response)]}126 127######################################## DOCUMENT WRITER AGENT ###################################3128                documenter_llm = ChatOpenAI(model="gpt-4o-mini")129 130                documenter_llm_prompt = ChatPromptTemplate.from_messages([131                        ("system", documenter_prompt),132                        ("human", "{content}")133                    ])134                135                documenter_chain = (136                    {"content": itemgetter("content")}137                    | documenter_llm_prompt138                    | documenter_llm139                    | StrOutputParser()140                )141 142                def write_document_content(state):143                    print(state)144                    json_content = state['messages'][-1].content145                    json_content = json_content[json_content.find("{"):json_content.rfind("}")+1].strip()146                    json_content = json.loads(json_content)147                    document_response = documenter_chain.invoke({"content": json_content})148                    return {"messages": [AIMessage(content=document_response)]}149                150########################################## CONSTRUCT GRAPH ############################################################33151                class AgentState(TypedDict):152                    messages: Annotated[list, add_messages]153 154                uncompiled_code_graph = StateGraph(AgentState)155                uncompiled_code_graph.add_node("code_agent", describe_code)156                uncompiled_code_graph.add_node("write_content_agent", write_document_content)157                uncompiled_code_graph.add_node("write_document", write_to_docx)158 159                uncompiled_code_graph.set_entry_point("code_agent")160                uncompiled_code_graph.add_edge("code_agent", "write_content_agent")161                uncompiled_code_graph.add_edge("write_content_agent", "write_document")162 163                compiled_code_graph = uncompiled_code_graph.compile()164 165 166                initial_state = {167                    "messages": [{168                        "role": "human",169                        "content": json.dumps({170                            "code_language": "python",171                            "code": file_content172                        })173                    }]174                }175                # bind model to tool or ToolNode176                # imports_tool_node = ToolNode(search_packages_tools)177 178                # construct graph and compile179                # uncompiled_imports_graph = StateGraph(AgentState)180                # uncompiled_imports_graph.add_node("imports_agent", call_imports_chain)181                # uncompiled_imports_graph.add_node("imports_action", imports_tool_node)182                # uncompiled_imports_graph.set_entry_point("imports_agent")183                184                # def should_continue(state):185                #     last_message = state["messages"][-1]186 187                #     if last_message.tool_calls:188                #         return "imports_action"189 190                #     return END191 192                # uncompiled_imports_graph.add_conditional_edges(193                #     "imports_agent",194                #     should_continue195                # )196 197                # uncompiled_imports_graph.add_edge("imports_action", "imports_agent")198 199                # compiled_imports_graph = uncompiled_imports_graph.compile()200 201                # print(f'compiled imports graph')202                # # Invoke imports graph203                # initial_state = {204                #     "messages": [{205                #         "role": "human",206                #         "content": json.dumps({207                #             "code_language": "python",208                #             "imports": imports209                #         })210                #     }]211                # }212 213 214 215 216 217                # await msg.update(content="Analyzing imports and generating documentation...")218                msg.content = "Analyzing your code and generating documentation..."219                await msg.update()220 221                # msg = cl.Message(content="Analyzing your code and generating documentation...")222                # await msg.send()223 224                documenter_result = compiled_code_graph.invoke(initial_state)225 226############################################## SAVE DESCRIPTION CHUNKS IN VECTOR STORE ########################################3227                qdrant_client = QdrantClient(":memory:")228 229                embedding_model = OpenAIEmbeddings(model="text-embedding-3-small")230                embedding_dim = 1536231 232                qdrant_client.create_collection(233                    collection_name="description_rag_data",234                    vectors_config=VectorParams(size=embedding_dim, distance=Distance.COSINE),235                )236 237                vectorstore = Qdrant(qdrant_client, collection_name="description_rag_data", embeddings=embedding_model)238 239                # Add chunks240                chunks = documenter_result['messages'][1].content241                chunks = chunks[chunks.find("{"):chunks.rfind("}")+1].strip()242                chunks = json.loads(chunks)243                print(f'################################### raw chunks \n {chunks} \n ######################## \n')244                chunks_list = []245                for key in chunks:246                    if isinstance(chunks[key], dict):247                        chunks_list.append(chunks[key])248                    elif isinstance(chunks[key], list):249                        for value in chunks[key]:250                            chunks_list.append(value)251                print(f'################################### chunks_list \n {chunks_list} \n ######################## \n')252                docs = [253                    Document(254                        page_content=f"{chunk.get('type', '')} - {chunk.get('name', '')} - {chunk.get('description', '')}",  # Content for the model255                        metadata={**chunk}  # Store metadata, but don't put embeddings here256                    )257                    for chunk in chunks_list258                ]259 260 261 262                vectorstore.add_documents(docs)263                qdrant_retriever = vectorstore.as_retriever(search_kwargs={"k": 3})264 265                print('done adding docs to DB')266                #define documenter chain267                # documenter_llm = ChatOpenAI(model="gpt-4o-mini")268                # documenter_llm_prompt = ChatPromptTemplate.from_messages([269                #     ("system", documenter_prompt),270                # ])271                # documenter_chain = (272                #     {"context": itemgetter("context")}273                #     | documenter_llm_prompt274                #     | documenter_llm275                #     | StrOutputParser()276                # )277 278                # print('done defining documenter chain')279 280                #extract description chunks from database281                # collection_name = "description_rag_data"282                # all_points = qdrant_client.scroll(collection_name=collection_name, limit=1000)[0]  # Adjust limit if needed283                # one_chunk = all_points[0].payload284                # input_text = f"type: {one_chunk['metadata']['type']} \nname: {one_chunk['metadata']['name']} \ncontent: {one_chunk['metadata']['content']}"285 286                # print('done extracting chunks form DB')287 288                # document_response = documenter_chain.invoke({"context": input_text})289 290                print('done invoking documenter chain and will write in docx')291                # write packages description in word file292                # document_file_path  = write_to_docx(document_response)293                # print (f'################################ \n documenter_result \n {documenter_result} \n ############################ \n')294                # document_file_path  = documenter_result['messages'][-1].content[0]295                # print()296                document_file_path = 'generated_documentation.docx'297 298 299                print('done writing docx file')300                # Set up Main Chain for chat301                main_llm = ChatOpenAI(model="gpt-4o-mini", temperature=0)302 303 304                main_llm_prompt = ChatPromptTemplate.from_messages([305                    ("system", main_prompt),306                    ("human", "{query}")307                ])308 309                main_chain = (310                    {"context": itemgetter("query") | qdrant_retriever, "code_language": itemgetter("code_language"), "query": itemgetter("query"), }311                    | main_llm_prompt312                    | main_llm313                    | StrOutputParser()314                )315 316                print('done defining main chain')317                # Present download button for the document318                elements = [319                    cl.File(320                        name="documentation.docx",321                        path=document_file_path,322                        display="inline"323                    )324                ]325                print('done defining elements')326                msg.content = "✅ Your Python file has been processed! You can download the documentation file below. How can I help you with your code?"327                msg.elements = elements328                await msg.update()329 330                # await msg.update(331                #         content="✅ Your Python file has been processed! You can download the documentation file below. How can I help you with your code?.",332                #         elements=elements333                #     )334                335            except Exception as e:336                # await msg.update(content=f"❌ Error processing file: {str(e)}")337                error_traceback = traceback.format_exc()338                print(error_traceback)339                msg.content = f"❌ Error processing file: {str(e)}"340                await msg.update()341 342                # msg = cl.Message(content=f"second message ❌ Error processing file: {str(e)}")343                # await msg.send()344        345        else:346            await cl.Message(content="Please upload a Python (.py) file.").send()347 348    # Handle chat messages if file has been processed349    elif processed_file_path and main_chain:350        user_input = message.content351        # Send thinking message352        msg = cl.Message(content="Thinking...")353        await msg.send()354 355        try:356            # Use main_chain to answer the query357# invoke main chain358            inputs = {359                'code_language': 'Python',360                'query': user_input361            }362 363            response = main_chain.invoke(inputs)364 365            # Update with the response366            # await msg.update(content=response)367            msg.content = response368            await msg.update()369 370            # msg = cl.Message(content=response)371            # await msg.send()372 373        except Exception as e:374            # await msg.update(content=f"❌ Error processing your question: {str(e)}")375            msg.content = f"❌ Error processing your question: {str(e)}"376            await msg.update()377 378            # msg = cl.Message(content=f"❌ Error processing your question: {str(e)}")379            # await msg.send()380 381    else:382        await cl.Message(content="Please upload a Python file first before asking questions.").send()383 384 385@cl.on_stop386def on_stop():387    global processed_file_path388    # Clean up temporary files389    if processed_file_path and os.path.exists(os.path.dirname(processed_file_path)):390        shutil.rmtree(os.path.dirname(processed_file_path))391 392