MaheshLEO4 commited on
Commit
fcbb0ab
·
1 Parent(s): 9bbaf64

updated files

Browse files
agents/research_agent.py CHANGED
@@ -1,52 +1,14 @@
1
- from typing import Dict, List
2
  from langchain.schema import Document
3
  from config.llm_config import llm_config
4
- import logging
5
-
6
- logger = logging.getLogger(__name__)
7
 
8
  class ResearchAgent:
9
  def __init__(self):
10
- """
11
- Initialize the research agent with LLM from LLMConfig
12
- """
13
- logger.info("Initializing ResearchAgent...")
14
-
15
- # Get LLM callable for research
16
- self.llm_fn = llm_config.create_llm("research")
17
-
18
- logger.info("ResearchAgent initialized successfully.")
19
-
20
- def generate(self, question: str, documents: List[Document]) -> Dict:
21
- """
22
- Generate an answer based on documents
23
- """
24
- logger.info(f"ResearchAgent.generate called with question='{question}' and {len(documents)} documents.")
25
-
26
- context = "\n\n".join([doc.page_content for doc in documents])
27
-
28
- prompt = f"""
29
- You are an AI assistant designed to provide precise and factual answers based on the given context.
30
-
31
- Question: {question}
32
-
33
- Context:
34
- {context}
35
-
36
- Provide your answer below:
37
- """
38
-
39
- try:
40
- draft_answer = self.llm_fn(prompt)
41
- logger.info(f"Generated answer successfully. Length: {len(draft_answer)} characters.")
42
- return {
43
- "draft_answer": draft_answer.strip(),
44
- "context_used": context
45
- }
46
-
47
- except Exception as e:
48
- logger.error(f"Error during answer generation: {e}")
49
- return {
50
- "draft_answer": f"I cannot answer this question based on the provided documents. Error: {str(e)}",
51
- "context_used": context
52
- }
 
 
1
  from langchain.schema import Document
2
  from config.llm_config import llm_config
 
 
 
3
 
4
  class ResearchAgent:
5
  def __init__(self):
6
+ self.llm = llm_config.create_llm("research")
7
+
8
+ def generate(self, question: str, documents: list[Document]):
9
+ context = "\n\n".join([d.page_content for d in documents])
10
+ prompt = f"Question: {question}\n\nContext: {context}\n\nAnswer:"
11
+
12
+ # Proper LangChain invocation
13
+ response = self.llm.invoke(prompt)
14
+ return {"draft_answer": response.content}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
agents/workflow.py CHANGED
@@ -1,11 +1,10 @@
1
- from typing import TypedDict, List, Dict
2
  from langchain.schema import Document
3
- from langchain.retrievers import EnsembleRetriever
 
4
  from .research_agent import ResearchAgent
5
  from .verification_agent import VerificationAgent
6
  from .relevance_checker import RelevanceChecker
7
- from langgraph.graph import StateGraph, END
8
- import logging
9
 
10
  logger = logging.getLogger(__name__)
11
 
@@ -15,89 +14,67 @@ class AgentState(TypedDict):
15
  draft_answer: str
16
  verification_report: str
17
  is_relevant: bool
18
- retriever: EnsembleRetriever
19
 
20
  class AgentWorkflow:
21
  def __init__(self):
22
  self.researcher = ResearchAgent()
23
  self.verifier = VerificationAgent()
24
  self.relevance_checker = RelevanceChecker()
25
- self.compiled_workflow = self.build_workflow()
26
 
27
  def build_workflow(self):
28
- workflow = StateGraph(AgentState)
29
- workflow.add_node("check_relevance", self._check_relevance_step)
30
- workflow.add_node("research", self._research_step)
31
- workflow.add_node("verify", self._verification_step)
 
32
 
33
- workflow.set_entry_point("check_relevance")
34
- workflow.add_conditional_edges(
 
35
  "check_relevance",
36
- self._decide_after_relevance_check,
37
- {"relevant": "research", "irrelevant": END}
38
  )
39
- workflow.add_edge("research", "verify")
40
- workflow.add_conditional_edges(
 
 
41
  "verify",
42
  self._decide_next_step,
43
  {"re_research": "research", "end": END}
44
  )
45
- return workflow.compile()
46
-
47
- def _check_relevance_step(self, state: AgentState) -> Dict:
48
- classification = self.relevance_checker.check(
49
- question=state["question"],
50
- retriever=state["retriever"],
51
- k=20
52
- )
53
-
54
- if classification in ["CAN_ANSWER", "PARTIAL"]:
55
- return {"is_relevant": True}
56
- return {
57
- "is_relevant": False,
58
- "draft_answer": "This question isn't related to the uploaded document(s)."
59
- }
60
-
61
- def _decide_after_relevance_check(self, state: AgentState) -> str:
62
- return "relevant" if state["is_relevant"] else "irrelevant"
63
-
64
- def full_pipeline(self, question: str, retriever: EnsembleRetriever):
65
- try:
66
- documents = retriever.get_relevant_documents(question) # updated method
67
- logger.info(f"Retrieved {len(documents)} documents")
68
 
69
- initial_state = AgentState(
70
- question=question,
71
- documents=documents,
72
- draft_answer="",
73
- verification_report="",
74
- is_relevant=False,
75
- retriever=retriever
76
- )
77
 
78
- final_state = self.compiled_workflow.invoke(initial_state)
79
- return {
80
- "draft_answer": final_state["draft_answer"],
81
- "verification_report": final_state["verification_report"]
82
- }
83
 
84
- except Exception as e:
85
- logger.error(f"Workflow execution failed: {e}")
86
- return {
87
- "draft_answer": "",
88
- "verification_report": f"Error: {e}"
89
- }
90
 
91
- def _research_step(self, state: AgentState) -> Dict:
92
- result = self.researcher.generate(state["question"], state["documents"])
93
- return {"draft_answer": result["draft_answer"]}
 
 
94
 
95
- def _verification_step(self, state: AgentState) -> Dict:
96
- result = self.verifier.check(state["draft_answer"], state["documents"])
97
- return {"verification_report": result["verification_report"]}
98
-
99
- def _decide_next_step(self, state: AgentState) -> str:
100
- report = state["verification_report"]
101
- if "NO" in report:
102
- return "re_research"
103
- return "end"
 
 
 
1
+ from typing import TypedDict, List, Dict, Annotated
2
  from langchain.schema import Document
3
+ from langgraph.graph import StateGraph, END
4
+ import logging
5
  from .research_agent import ResearchAgent
6
  from .verification_agent import VerificationAgent
7
  from .relevance_checker import RelevanceChecker
 
 
8
 
9
  logger = logging.getLogger(__name__)
10
 
 
14
  draft_answer: str
15
  verification_report: str
16
  is_relevant: bool
17
+ retry_count: int # Added to prevent infinite loops
18
 
19
  class AgentWorkflow:
20
  def __init__(self):
21
  self.researcher = ResearchAgent()
22
  self.verifier = VerificationAgent()
23
  self.relevance_checker = RelevanceChecker()
24
+ self.workflow = self.build_workflow()
25
 
26
  def build_workflow(self):
27
+ builder = StateGraph(AgentState)
28
+
29
+ builder.add_node("check_relevance", self._check_relevance_step)
30
+ builder.add_node("research", self._research_step)
31
+ builder.add_node("verify", self._verification_step)
32
 
33
+ builder.set_entry_point("check_relevance")
34
+
35
+ builder.add_conditional_edges(
36
  "check_relevance",
37
+ lambda x: "research" if x["is_relevant"] else "end",
38
+ {"research": "research", "end": END}
39
  )
40
+
41
+ builder.add_edge("research", "verify")
42
+
43
+ builder.add_conditional_edges(
44
  "verify",
45
  self._decide_next_step,
46
  {"re_research": "research", "end": END}
47
  )
48
+
49
+ return builder.compile()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
50
 
51
+ def _check_relevance_step(self, state: AgentState):
52
+ # Logic to call relevance_checker.check
53
+ res = self.relevance_checker.check(state["question"], state["documents"])
54
+ return {"is_relevant": res != "NO_MATCH", "retry_count": 0}
 
 
 
 
55
 
56
+ def _research_step(self, state: AgentState):
57
+ res = self.researcher.generate(state["question"], state["documents"])
58
+ return {"draft_answer": res["draft_answer"], "retry_count": state.get("retry_count", 0) + 1}
 
 
59
 
60
+ def _verification_step(self, state: AgentState):
61
+ res = self.verifier.check(state["draft_answer"], state["documents"])
62
+ return {"verification_report": res["verification_report"]}
 
 
 
63
 
64
+ def _decide_next_step(self, state: AgentState):
65
+ # Break loop after 2 retries or if supported
66
+ if "Supported: YES" in state["verification_report"] or state["retry_count"] >= 3:
67
+ return "end"
68
+ return "re_research"
69
 
70
+ def full_pipeline(self, question: str, retriever):
71
+ docs = retriever.invoke(question) # Updated from get_relevant_documents
72
+ initial_state = {
73
+ "question": question,
74
+ "documents": docs,
75
+ "draft_answer": "",
76
+ "verification_report": "",
77
+ "is_relevant": False,
78
+ "retry_count": 0
79
+ }
80
+ return self.workflow.invoke(initial_state)
config/llm_config.py CHANGED
@@ -1,113 +1,44 @@
1
- """
2
- LLM Configuration Manager
3
- Centralizes all LLM model configurations for easy switching
4
- """
5
- from typing import Dict, Any
6
- from enum import Enum
7
  import os
8
  import logging
9
-
10
- # Modern Google Generative AI package import
11
  import google.generativeai as genai
12
-
13
- # Environment variables
 
14
  from dotenv import load_dotenv
15
- load_dotenv()
16
 
 
17
  logger = logging.getLogger(__name__)
18
 
19
- class ModelProvider(Enum):
20
- GOOGLE = "google"
21
- OPENAI = "openai"
22
-
23
  class LLMConfig:
24
- """Configuration manager for LLM models"""
25
-
26
- MODELS = {
27
- ModelProvider.GOOGLE: {
 
 
 
 
 
28
  "research": "gemini-1.5-pro",
29
  "verification": "gemini-1.5-flash",
30
- "relevance": "gemini-1.5-flash",
31
- "embedding": "text-embedding-004",
32
- },
33
- ModelProvider.OPENAI: {
34
- "research": "gpt-4-turbo",
35
- "verification": "gpt-4-turbo",
36
- "relevance": "gpt-4-turbo",
37
- "embedding": "text-embedding-3-large",
38
  }
39
- }
40
-
41
- DEFAULT_PARAMS = {
42
- "research": {"temperature": 0.3, "max_tokens": 300, "top_p": 0.95},
43
- "verification": {"temperature": 0.0, "max_tokens": 200, "top_p": 0.9},
44
- "relevance": {"temperature": 0.0, "max_tokens": 10, "top_p": 0.9}
45
- }
46
-
47
- def __init__(self, provider: ModelProvider = ModelProvider.GOOGLE):
48
- self.provider = provider
49
- self.api_key = self._get_api_key()
50
- self._validate_config()
51
- genai.api_key = self.api_key
52
-
53
- def _get_api_key(self) -> str:
54
- if self.provider == ModelProvider.GOOGLE:
55
- key = os.getenv("GOOGLE_API_KEY")
56
- if not key:
57
- raise ValueError("GOOGLE_API_KEY environment variable is required")
58
- return key
59
- elif self.provider == ModelProvider.OPENAI:
60
- key = os.getenv("OPENAI_API_KEY")
61
- if not key:
62
- raise ValueError("OPENAI_API_KEY environment variable is required")
63
- return key
64
- raise ValueError(f"Unsupported provider: {self.provider}")
65
-
66
- def _validate_config(self):
67
- if self.provider not in self.MODELS:
68
- raise ValueError(f"Provider {self.provider} not configured")
69
-
70
- def get_model_name(self, task: str) -> str:
71
- return self.MODELS[self.provider][task]
72
-
73
- def get_model_params(self, task: str) -> Dict[str, Any]:
74
- return self.DEFAULT_PARAMS.get(task, {}).copy()
75
-
76
- def create_llm(self, task: str):
77
- """Return a callable that sends prompt to LLM and returns the response text."""
78
- model_name = self.get_model_name(task)
79
- params = self.get_model_params(task)
80
 
81
- def llm_callable(prompt: str) -> str:
82
- try:
83
- response = genai.chat.create(
84
- model=model_name,
85
- messages=[{"role": "user", "content": prompt}],
86
- temperature=params.get("temperature", 0.3),
87
- max_output_tokens=params.get("max_tokens", 300),
88
- top_p=params.get("top_p", 0.95),
89
- )
90
- return response.last.split("\n")[0] if hasattr(response, 'last') else response.choices[0].content
91
- except Exception as e:
92
- logger.error(f"LLM call failed: {e}")
93
- raise
94
- return llm_callable
95
-
96
  def create_embedding(self):
97
- """Return a callable for embeddings"""
98
- if self.provider == ModelProvider.GOOGLE:
99
- model_name = self.get_model_name("embedding")
100
- def embed_fn(texts: list[str]) -> list[list[float]]:
101
- try:
102
- response = genai.embeddings.create(
103
- model=model_name,
104
- input=texts
105
- )
106
- return [item.embedding for item in response.data]
107
- except Exception as e:
108
- logger.error(f"Embedding generation failed: {e}")
109
- raise
110
- return embed_fn
111
 
112
- # Global configuration instance
113
- llm_config = LLMConfig()
 
 
 
 
 
 
 
1
  import os
2
  import logging
 
 
3
  import google.generativeai as genai
4
+ from typing import List, Dict, Any
5
+ from langchain_core.embeddings import Embeddings
6
+ from langchain_google_genai import ChatGoogleGenerativeAI, GoogleGenerativeAIEmbeddings
7
  from dotenv import load_dotenv
 
8
 
9
+ load_dotenv()
10
  logger = logging.getLogger(__name__)
11
 
 
 
 
 
12
  class LLMConfig:
13
+ def __init__(self):
14
+ self.google_api_key = os.getenv("GOOGLE_API_KEY")
15
+ if not self.google_api_key:
16
+ raise ValueError("GOOGLE_API_KEY not found")
17
+ genai.configure(api_key=self.google_api_key)
18
+
19
+ def create_llm(self, task: str):
20
+ # Map tasks to models
21
+ model_map = {
22
  "research": "gemini-1.5-pro",
23
  "verification": "gemini-1.5-flash",
24
+ "relevance": "gemini-1.5-flash"
 
 
 
 
 
 
 
25
  }
26
+ params = {
27
+ "research": {"temperature": 0.3},
28
+ "verification": {"temperature": 0.0},
29
+ "relevance": {"temperature": 0.0}
30
+ }
31
+
32
+ return ChatGoogleGenerativeAI(
33
+ model=model_map[task],
34
+ google_api_key=self.google_api_key,
35
+ temperature=params[task]["temperature"]
36
+ )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
37
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
38
  def create_embedding(self):
39
+ return GoogleGenerativeAIEmbeddings(
40
+ model="models/text-embedding-004",
41
+ google_api_key=self.google_api_key
42
+ )
 
 
 
 
 
 
 
 
 
 
43
 
44
+ llm_config = LLMConfig()
 
requirements.txt CHANGED
@@ -1,3 +1,12 @@
 
 
 
 
 
 
 
 
 
1
  # Core Python
2
  python-dotenv
3
  pydantic
 
1
+ gradio
2
+ langchain
3
+ langchain-google-genai
4
+ langgraph
5
+ chromadb
6
+ pydantic-settings
7
+ python-dotenv
8
+ pypdf
9
+ docx2txt
10
  # Core Python
11
  python-dotenv
12
  pydantic
retriever/builder.py CHANGED
@@ -9,47 +9,22 @@ logger = logging.getLogger(__name__)
9
 
10
  class RetrieverBuilder:
11
  def __init__(self):
12
- """Initialize the retriever builder with embeddings."""
13
- logger.info("Initializing RetrieverBuilder...")
14
-
15
- # Get embeddings from configuration
16
  self.embeddings = llm_config.create_embedding()
17
 
18
- logger.info("RetrieverBuilder initialized successfully.")
19
-
20
  def build_hybrid_retriever(self, docs):
21
- """Build a hybrid retriever using BM25 and vector-based retrieval."""
22
- try:
23
- logger.info(f"Building hybrid retriever with {len(docs)} documents")
24
-
25
- # Create Chroma vector store
26
- vector_store = Chroma.from_documents(
27
- documents=docs,
28
- embedding=self.embeddings,
29
- persist_directory=settings.CHROMA_DB_PATH,
30
- collection_name=settings.CHROMA_COLLECTION_NAME
31
- )
32
- logger.info("Vector store created successfully.")
33
-
34
- # Create BM25 retriever
35
- bm25 = BM25Retriever.from_documents(docs)
36
- logger.info("BM25 retriever created successfully.")
37
-
38
- # Create vector-based retriever
39
- vector_retriever = vector_store.as_retriever(
40
- search_kwargs={"k": settings.VECTOR_SEARCH_K}
41
- )
42
- logger.info("Vector retriever created successfully.")
43
-
44
- # Combine retrievers into a hybrid retriever
45
- hybrid_retriever = EnsembleRetriever(
46
- retrievers=[bm25, vector_retriever],
47
- weights=settings.HYBRID_RETRIEVER_WEIGHTS
48
- )
49
- logger.info("Hybrid retriever created successfully.")
50
-
51
- return hybrid_retriever
52
-
53
- except Exception as e:
54
- logger.error(f"Failed to build hybrid retriever: {e}")
55
- raise
 
9
 
10
  class RetrieverBuilder:
11
  def __init__(self):
12
+ # Correctly get the LangChain Embedding Object
 
 
 
13
  self.embeddings = llm_config.create_embedding()
14
 
 
 
15
  def build_hybrid_retriever(self, docs):
16
+ # Use the class-based embedding interface
17
+ vector_store = Chroma.from_documents(
18
+ documents=docs,
19
+ embedding=self.embeddings,
20
+ persist_directory=settings.CHROMA_DB_PATH,
21
+ collection_name=settings.CHROMA_COLLECTION_NAME
22
+ )
23
+
24
+ bm25 = BM25Retriever.from_documents(docs)
25
+ vector_retriever = vector_store.as_retriever(search_kwargs={"k": 5})
26
+
27
+ return EnsembleRetriever(
28
+ retrievers=[bm25, vector_retriever],
29
+ weights=[0.4, 0.6]
30
+ )