Team Ai
Apppublic

melghorab/code-assistant

sourceHugging Faceopenrailupdated 2y agoView on Hugging Face
0likes
app.py318 linesDownload Raw Back to root
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