11"""Chat endpoint — submits a query to the LangGraph RAG agent."""
22
3+ import asyncio
34import json
45import time
56
6- from fastapi import APIRouter , Depends
7+ from fastapi import APIRouter , Depends , HTTPException , status
78from langchain_core .messages import HumanMessage , ToolMessage
9+ from langgraph .errors import GraphRecursionError
810
911from app .agent .graph import app_graph
10- from app .api .dependencies import require_company_user
12+ from app .agent .tools .vector_search import (
13+ reset_search_company_scope ,
14+ reset_search_query_state ,
15+ set_search_company_scope ,
16+ set_search_query_state ,
17+ warmup_vector_search ,
18+ )
19+ from app .api .dependencies import get_current_user , require_company_user
1120from app .core .exceptions import AgentError
1221from app .core .logger import get_logger
1322from app .db .models import User
14- from app .schemas .chat import ChatRequest , ChatResponse
23+ from app .schemas .chat import ChatRequest , ChatResponse , WarmupResponse
1524
1625logger = get_logger (__name__ )
1726router = APIRouter ()
1827
1928
29+ @router .post (
30+ "/warmup" ,
31+ response_model = WarmupResponse ,
32+ responses = {
33+ 502 : {"description" : "Warmup failed — embedding model could not be loaded." },
34+ },
35+ )
36+ async def warmup_agent (current_user : User = Depends (get_current_user )):
37+ """Warm up agent resources so the next chat request avoids embedding-model cold start latency."""
38+ start = time .monotonic ()
39+ try :
40+ loaded_now = await asyncio .to_thread (warmup_vector_search )
41+ except Exception as exc :
42+ elapsed = time .monotonic () - start
43+ logger .error (
44+ "Agent warmup failed" ,
45+ extra = {"user_id" : current_user .id , "elapsed_s" : round (elapsed , 3 )},
46+ exc_info = exc ,
47+ )
48+ raise AgentError ("Agent warmup failed. Please try again." ) from exc
49+
50+ elapsed = time .monotonic () - start
51+ logger .info (
52+ "Agent warmup completed" ,
53+ extra = {
54+ "user_id" : current_user .id ,
55+ "elapsed_s" : round (elapsed , 3 ),
56+ "embeddings_loaded_now" : loaded_now ,
57+ },
58+ )
59+ return WarmupResponse (
60+ message = "Agent warmup completed" ,
61+ embeddings_loaded_now = loaded_now ,
62+ elapsed_seconds = round (elapsed , 3 ),
63+ )
64+
65+
2066def _extract_graph_payload (messages ) -> dict | None :
2167 for msg in reversed (messages ):
2268 if isinstance (msg , ToolMessage ) and msg .name == "generate_graph" :
@@ -61,10 +107,36 @@ async def invoke_agent(
61107 },
62108 )
63109
110+ if not current_user .company_id :
111+ logger .warning (
112+ "Chat blocked — user has no company scope" ,
113+ extra = {"user_id" : current_user .id , "role" : str (current_user .role )},
114+ )
115+ raise HTTPException (
116+ status_code = status .HTTP_403_FORBIDDEN ,
117+ detail = "You do not have permission to perform this action." ,
118+ )
119+
64120 start = time .monotonic ()
121+ search_scope_token = set_search_company_scope (current_user .company_id )
122+ search_query_state_token = set_search_query_state ()
65123 try :
66124 initial_state = {"messages" : [HumanMessage (content = request .query )]}
67- final_state = await app_graph .ainvoke (initial_state )
125+ final_state = await app_graph .ainvoke (initial_state , config = {"recursion_limit" : 8 })
126+ except GraphRecursionError as exc :
127+ elapsed = time .monotonic () - start
128+ logger .error (
129+ "Agent halted by recursion limit" ,
130+ extra = {
131+ "user_id" : current_user .id ,
132+ "company_id" : current_user .company_id ,
133+ "elapsed_s" : round (elapsed , 3 ),
134+ },
135+ exc_info = exc ,
136+ )
137+ raise AgentError (
138+ "The agent could not converge on an answer. Please rephrase your question with more specific details."
139+ ) from exc
68140 except Exception as exc :
69141 elapsed = time .monotonic () - start
70142 logger .error (
@@ -77,6 +149,9 @@ async def invoke_agent(
77149 exc_info = exc ,
78150 )
79151 raise AgentError ("The agent failed to process your request. Please try again." ) from exc
152+ finally :
153+ reset_search_query_state (search_query_state_token )
154+ reset_search_company_scope (search_scope_token )
80155
81156 elapsed = time .monotonic () - start
82157 final_message = final_state ["messages" ][- 1 ].content
0 commit comments