Team Ai
Apppublic

kyorlin2001/code-analysis-tool

sourceHugging Faceapache-2.0updated 6mo agoView on Hugging Face
0likes
rag_agent.py145 linesDownload Raw Back to rag
1from __future__ import annotations2 3from dataclasses import dataclass, field4from typing import Any5 6from config.model_config import ModelConfig7from models.rag_result import RagResult8from rag.answer_merger import AnswerMerger9from rag.citation_formatter import CitationFormatter10from rag.context_budget import ContextBudgetManager11from rag.model_client import ModelClient12from rag.prompt_builder import PromptBuilder13from rag.retriever import Retriever14 15 16@dataclass17class RagAgentInput:18    """19    Input payload for a RAG query.20    """21 22    question: str23    repo_name: str | None = None24    findings: list[dict[str, Any]] = field(default_factory=list)25    top_k: int | None = None26 27 28class RagAgent:29    """30    Coordinates retrieval and model reasoning for repository questions.31    """32 33    def __init__(34        self,35        retriever: Retriever,36        model_client: ModelClient | None = None,37        prompt_builder: PromptBuilder | None = None,38        context_budget_manager: ContextBudgetManager | None = None,39        citation_formatter: CitationFormatter | None = None,40        answer_merger: AnswerMerger | None = None,41        config: ModelConfig | None = None,42    ) -> None:43        self.retriever = retriever44        self.config = config or ModelConfig.from_env()45        self.model_client = model_client or ModelClient(self.config)46        self.prompt_builder = prompt_builder or PromptBuilder()47        self.context_budget_manager = context_budget_manager or ContextBudgetManager(48            max_context_chars=self.config.max_context_chars49        )50        self.citation_formatter = citation_formatter or CitationFormatter()51        self.answer_merger = answer_merger or AnswerMerger()52 53    def run(self, payload: RagAgentInput) -> RagResult:54        top_k = payload.top_k or self.config.top_k55        retrieval = self.retriever.retrieve(payload.question, top_k=top_k)56 57        budget_result = self.context_budget_manager.apply(retrieval.chunks)58        selected_chunks = budget_result.selected_chunks59 60        prompt = self.prompt_builder.build(61            question=payload.question,62            chunks=selected_chunks,63            repo_name=payload.repo_name,64            findings=payload.findings,65        )66 67        model_response = self.model_client.complete(prompt)68        parsed = self._parse_response_text(model_response.text)69 70        citations = [71            self.citation_formatter.format_chunk(chunk).__dict__72            for chunk in selected_chunks73        ]74 75        debug_info = {76            "retrieved_count": len(retrieval.chunks),77            "selected_count": len(selected_chunks),78            "selected_files": [chunk.file_path for chunk in selected_chunks],79            "selected_chunk_previews": [chunk.text[:500] for chunk in selected_chunks[:5]],80            "budget_truncated": budget_result.truncated,81            "budget_total_characters": budget_result.total_characters,82            "prompt_preview": prompt.user_prompt[:4000],83        }84 85        return RagResult(86            answer=parsed["answer"],87            suggestions=parsed["suggestions"],88            citations=citations,89            follow_up_questions=parsed["follow_up_questions"],90            notes=parsed["notes"],91            raw_response={92                "model_response": model_response.raw_response,93                "debug_info": debug_info,94            },95        )96 97       # return rag_result98 99    def _parse_response_text(self, text: str) -> dict[str, Any]:100        """101        Parse model output into a structured result.102 103        Expected format:104        - answer: ...105        - suggestions:106          - ...107        - citations:108          - ...109        - follow_up_questions:110          - ...111        - notes:112          - ...113        """114        sections: dict[str, list[str]] = {115            "answer": [],116            "suggestions": [],117            "citations": [],118            "follow_up_questions": [],119            "notes": [],120        }121 122        current_section = "answer"123 124        for raw_line in text.splitlines():125            line = raw_line.strip()126            if not line:127                continue128 129            lowered = line.lower().rstrip(":")130            if lowered in sections:131                current_section = lowered132                continue133 134            if line.startswith("- "):135                sections[current_section].append(line[2:].strip())136            else:137                sections[current_section].append(line)138 139        return {140            "answer": "\n".join(sections["answer"]).strip(),141            "suggestions": sections["suggestions"],142            "citations": sections["citations"],143            "follow_up_questions": sections["follow_up_questions"],144            "notes": sections["notes"],145        }