import asyncio from fastapi import APIRouter, HTTPException, Depends from pydantic import BaseModel from typing import List, Optional # pyrefly: ignore [missing-import] from langchain_core.messages import HumanMessage from app.agents.graph import app_graph from app.api.deps import get_current_user from app.core.config import settings from huggingface_hub import HfApi router = APIRouter() class ChatRequest(BaseModel): message: str session_id: Optional[str] = None chat_history: Optional[List[dict]] = None class ChatResponse(BaseModel): reply: str route_taken: Optional[str] = None from fastapi.responses import StreamingResponse import json @router.post("/chat/stream") async def chat_stream_endpoint(request: ChatRequest, current_user_id: str = Depends(get_current_user)): # Emergency Kill Switch if settings.EMERGENCY_KILL_COMMAND and request.message.strip() == settings.EMERGENCY_KILL_COMMAND: async def kill_generator(): yield f"data: {json.dumps({'chunk': '🚨 EMERGENCY STOP ACTIVATED. Pausing Hugging Face Space...'})}\n\n" yield f"data: {json.dumps({'route_taken': 'system'})}\n\n" try: # Pause the space asynchronously to allow the message to stream first import threading def pause_space(): try: api = HfApi(token=settings.HF_TOKEN) api.pause_space(repo_id=settings.HF_REPO_ID) except Exception as e: print(f"Failed to pause space: {e}") threading.Thread(target=pause_space).start() except Exception: pass return StreamingResponse(kill_generator(), media_type="text/event-stream") input_text = request.message # Build messages list from history msgs = [] if request.chat_history: for m in request.chat_history: if m.get("role") == "user": msgs.append(HumanMessage(content=m.get("content", ""))) else: from langchain_core.messages import AIMessage msgs.append(AIMessage(content=m.get("content", ""))) msgs.append(HumanMessage(content=input_text)) initial_state = { "messages": msgs, "next_node": "" } async def event_generator(): try: final_reply = "" route_taken = "general_agent" # default fallback final_state = None async for event in app_graph.astream_events(initial_state, version="v2"): kind = event["event"] name = event.get("name", "") # Identify which tool was called to determine the route if kind == "on_tool_start": if name == "campus_data": route_taken = "academic_agent" data_str = json.dumps({'status': 'Querying Campus Knowledge...'}) yield f"data: {data_str}\n\n" elif name == "latest_announcements": route_taken = "campus_agent" data_str = json.dumps({'status': 'Checking Latest Announcements...'}) yield f"data: {data_str}\n\n" elif name == "google_search_tool": route_taken = "general_agent" data_str = json.dumps({'status': 'Searching the Web...'}) yield f"data: {data_str}\n\n" if kind == "on_chat_model_stream": chunk = event["data"]["chunk"].content if isinstance(chunk, list): chunk_text = "" for item in chunk: if isinstance(item, dict) and "text" in item: chunk_text += item["text"] elif isinstance(item, str): chunk_text += item chunk = chunk_text if chunk: final_reply += chunk yield f"data: {json.dumps({'chunk': chunk})}\n\n" if kind == "on_chain_end": output = event.get("data", {}).get("output") if isinstance(output, dict) and "messages" in output: final_state = output # If the LLM didn't stream anything (e.g., due to hitting recursion limit and returning a fallback) if not final_reply and final_state and "messages" in final_state: last_msg = final_state["messages"][-1] if getattr(last_msg, "type", "") == "ai" and getattr(last_msg, "content", ""): yield f"data: {json.dumps({'chunk': last_msg.content})}\n\n" # Yield the final route metadata yield f"data: {json.dumps({'route_taken': route_taken})}\n\n" except Exception as e: print(f"Error in chat stream: {str(e)}") # Log the real error for backend debugging user_friendly_error = "Error Sending message :\n1) Check your internet connection.\n2) We might be experiencing high traffic. Please try again later." yield f"data: {json.dumps({'error': user_friendly_error})}\n\n" return StreamingResponse(event_generator(), media_type="text/event-stream") @router.post("/chat", response_model=ChatResponse) async def chat_endpoint(request: ChatRequest, current_user_id: str = Depends(get_current_user)): # Emergency Kill Switch if settings.EMERGENCY_KILL_COMMAND and request.message.strip() == settings.EMERGENCY_KILL_COMMAND: import threading def pause_space(): try: api = HfApi(token=settings.HF_TOKEN) api.pause_space(repo_id=settings.HF_REPO_ID) except Exception as e: print(f"Failed to pause space: {e}") threading.Thread(target=pause_space).start() return ChatResponse(reply="🚨 EMERGENCY STOP ACTIVATED. Pausing Hugging Face Space...", route_taken="system") # Fallback non-streaming endpoint input_text = request.message # Build messages list from history msgs = [] if request.chat_history: for m in request.chat_history: if m.get("role") == "user": msgs.append(HumanMessage(content=m.get("content", ""))) else: from langchain_core.messages import AIMessage msgs.append(AIMessage(content=m.get("content", ""))) msgs.append(HumanMessage(content=input_text)) initial_state = { "messages": msgs, "next_node": "" } # Run the graph asynchronously try: final_state = await app_graph.ainvoke(initial_state) # The reply is the last message in the list reply = final_state["messages"][-1].content if isinstance(reply, list): reply_text = "" for item in reply: if isinstance(item, dict) and "text" in item: reply_text += item["text"] elif isinstance(item, str): reply_text += item reply = reply_text return ChatResponse( reply=reply, route_taken=final_state.get("next_node") ) except Exception as e: print(f"Error in chat endpoint: {str(e)}") # Log the real error for backend debugging user_friendly_error = "Error Sending message :\n1) Check your internet connection.\n2) We might be experiencing high traffic. Please try again later." raise HTTPException(status_code=500, detail=user_friendly_error)