import asyncio import logging import os from contextlib import asynccontextmanager from pathlib import Path from dotenv import load_dotenv from ecologits import EcoLogits from fastapi import BackgroundTasks, FastAPI, File, Form, Request, Response, UploadFile from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import HTMLResponse from fastapi.staticfiles import StaticFiles from fastapi.templating import Jinja2Templates from google import genai from huggingface_hub import InferenceClient from openai import OpenAI from pillow_heif import register_heif_opener from slowapi import Limiter from slowapi.util import get_remote_address from uvicorn.logging import DefaultFormatter # Important to import ecologits_patch before ecologits from agent.agent import Agent from agent.skill_manager import SkillsManager from classes.base_models import ( ChatRequest, CommentRequest, DeleteFileRequest, FeedbackRequest, ) from classes.session_document_store import SessionDocumentStore from classes.session_router import SessionRouter from classes.session_tracker import SessionTracker from constants import ( GEMINI_MODEL, GENERIC_ERROR_REPLY, HF_TOKEN, MAX_ID_LENGTH, OPENAI_MODEL, STATUS_CODE_INTERNAL_SERVER_ERROR, ) from exceptions import ( FILE_DOCUMENT_STORE_ERROR_STATUS_CODES, FILE_EXTRACTION_ERROR_STATUS_CODES, FILE_VALIDATION_ERROR_STATUS_CODES, FileDocumentStoreException, FileExtractionException, FileValidationException, ) from helpers.dynamodb_helper import ( init_dynamodb_tables, log_chat_event, log_environment_event, ) from helpers.file_helper import ( extract_text_from_file, replace_spaces_in_filename, validate_file, ) from helpers.lifespan_helper import cleanup_loop, load_heavy_models, run_cleanup from helpers.rag import create_embedding_model, load_vector_store from helpers.timing import request_timings from providers.fake import FakeProvider from providers.gemini import GeminiProvider from providers.hf import HFChatProvider from providers.openai import OpenAIProvider from telemetry import setup_telemetry load_dotenv() logger = logging.getLogger("uvicorn") # -------------------- Config -------------------- DEV = os.getenv("ENV", None) == "dev" logger.info(f"OPENAI_MODEL: {OPENAI_MODEL}") logger.info(f"GEMINI_MODEL: {GEMINI_MODEL}") # -------------------- Helpers -------------------- # For now, conversations and uploaded documents are stored in RAM. # This is tolerable for a demo, but we will have to switch to # Redis (or another real-time database) at some point. We are # currently storing sessions in what should be a stateless server. session_tracker = SessionTracker() session_document_store = SessionDocumentStore() # -------------------- Environmental Impact -------------------- tracker = None # tracker = EmissionsTracker( # project_name="test", measure_power_secs=5, save_to_file=False # ) # tracker.start() if tracker is not None: logger.info(f"Detected hardware: {tracker.get_detected_hardware()}") logger.info(f"Geographic metadata: {tracker._geo}") else: logger.info("Infrastructure carbon emissions tracking is disabled.") def log_environment_infra(): gwp_emissions = tracker.flush() try: infra_data = { "energy_kWh": tracker._total_energy.kWh, "co2eq_kg": gwp_emissions, "water_L": tracker._total_water.litres, } log_environment_event("infrastructure", infra_data) except Exception as e: logger.error(e) async def environment_infra_loop(): """Background task that runs forever while the app is alive.""" while True: await asyncio.sleep(3600) # 1 hour log_environment_infra() # -------------------- FastAPI setup -------------------- @asynccontextmanager async def lifespan(app: FastAPI): # Setup logging logger = logging.getLogger("uvicorn") if logger.handlers: colored_formatter = DefaultFormatter( fmt="%(levelprefix)s %(asctime)s | %(message)s", datefmt="%Y-%m-%d %H:%M:%S" ) logger.handlers[0].setFormatter(colored_formatter) logger.info("Logging configured!") # Setup heavy models load_heavy_models() # Connect/create DynamoDB tables (chat + environment logging) init_dynamodb_tables() # Setup Ecologits EcoLogits.init( providers=["huggingface_hub", "openai", "google_genai"], electricity_mix_zone="USA", ) # Setup CodeCarbon environment_infra_bg_task = None if tracker is not None: environment_infra_bg_task = asyncio.create_task(environment_infra_loop()) # Setup cleanup loop cleanup_bg_task = asyncio.create_task( cleanup_loop( session_tracker, session_document_store, session_router, ) ) # Register HEIF opener so Pillow can support HEIF/HEIC files (file upload features) register_heif_opener() yield cleanup_bg_task.cancel() if environment_infra_bg_task is not None: environment_infra_bg_task.cancel() app = FastAPI(lifespan=lifespan) setup_telemetry(app) if DEV: # Dev-only: lets a locally-run frontend (or a teammate's machine on the # same LAN) hit this API cross-origin. Never enabled in the deployed app. app.add_middleware( CORSMiddleware, allow_origins=[ "http://localhost:8081", "http://127.0.0.1:8081", "http://localhost:8000", ], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) app.mount("/static", StaticFiles(directory="static"), name="static") templates = Jinja2Templates(directory="templates") @app.middleware("http") async def cleanup_middleware(request: Request, call_next): run_cleanup( session_tracker, session_document_store, session_router, ) response = await call_next(request) return response @app.get("/", response_class=HTMLResponse) async def home(request: Request): return templates.TemplateResponse( name="index.html", request=request, context={"dev": DEV} ) # Rate limiter limiter = Limiter(key_func=get_remote_address) skills_dir = Path().cwd() / "agent/skills" client = InferenceClient(api_key=HF_TOKEN, provider="groq") provider = HFChatProvider(client) # Production pediatric agent: grounded on the Karpathy-style wiki via the # `pediatry_wiki_minimal` skill (MINIMAL_ANSWER_PROMPT draft, MINIMAL_JUDGE_PROMPT # judge). `pediatry_wiki` — same pipeline/wiki, production JUDGE_PROMPT and # GENERATE_ANSWER_PROMPT_SHORT instead — is kept as the pre-switch baseline, # not served. The other pediatry_wiki_* variants (short, self_judge, verbatim, # reject_ledger, reject_ledger_split_index, no_critic) were eval-only # comparison baselines and have been retired. _COMMON_SKILLS_TO_IGNORE = [ "basic_hiv_facts", "calculate", "meds_identification", # Baseline kept for comparison against pediatry_wiki_minimal — not served. "pediatry_wiki", ] wiki_skills = SkillsManager( skills_dir=skills_dir, skills_to_ignore=_COMMON_SKILLS_TO_IGNORE, ) wiki_skills.discover() wiki_skills_agent = Agent(wiki_skills, provider) champ_vector_store = load_vector_store(create_embedding_model()) openai_provider = OpenAIProvider(OpenAI(api_key=os.getenv("OPENAI_API_KEY"))) gemini_provider = GeminiProvider(genai.Client(api_key=os.getenv("GEMINI_API_KEY"))) fake_provider = FakeProvider() session_router = SessionRouter( wiki_skills_agent=wiki_skills_agent, champ_provider=provider, champ_vector_store=champ_vector_store, openai_provider=openai_provider, openai_model_id=OPENAI_MODEL, gemini_provider=gemini_provider, gemini_model_id=GEMINI_MODEL, fake_provider=fake_provider, ) @app.post("/chat") @limiter.limit("450/minute") def chat_endpoint( payload: ChatRequest, background_tasks: BackgroundTasks, request: Request ): # Collects every timed_block fired down the call stack (agent loop, wiki # sub-agent pipeline) into one JSONL row per request — see helpers/timing.py. with request_timings( model_type=payload.model_type, session_id=payload.session_id, conversation_id=payload.conversation_id, lang=payload.lang, ): session_tracker.update_session(payload.session_id) documents = session_document_store.get_documents(payload.session_id) outcome = session_router.send(payload, documents=documents) chat_data: dict = { "model_type": payload.model_type, "consent": payload.consent, "human_message": payload.human_message, "age_group": payload.age_group, "gender": payload.gender, "roles": payload.roles, "participant_id": payload.participant_id, "conversation_id": payload.conversation_id, "lang": payload.lang, **(outcome.triage_meta or {}), } if outcome.success: chat_data.update( { "reply": outcome.reply, "reply_id": outcome.reply_id, "context": outcome.context, } ) else: chat_data["error"] = outcome.error # Per-turn internals (tool calls, sub-agent pipeline, judge verdicts, # reasoning) — attached on failures too, where it matters most. if outcome.trace: chat_data["trace"] = outcome.trace background_tasks.add_task( log_chat_event, user_id=payload.user_id, session_id=payload.session_id, data=chat_data, ) if outcome.success and outcome.inference_impacts is not None: background_tasks.add_task( log_environment_event, source_type="inference", data_obj=outcome.inference_impacts, model_type=payload.model_type, reply_id=outcome.reply_id, participant_id=payload.participant_id, ) return { "reply": outcome.reply if outcome.success else GENERIC_ERROR_REPLY, "reply_id": outcome.reply_id, "gwp_kgcoeq": outcome.env_impact.gwp_kgcoeq, "water_L": outcome.env_impact.water_L, "electricity_kWh": outcome.env_impact.electricity_kWh, "n_tokens": outcome.n_tokens, } # Endpoint for specific replies/responses @app.post("/feedback") @limiter.limit("450/minute") def feedback_endpoint( payload: FeedbackRequest, background_tasks: BackgroundTasks, request: Request ): background_tasks.add_task( log_chat_event, user_id=payload.user_id, session_id=payload.session_id, data={ "consent": payload.consent, "comment": payload.comment, "age_group": payload.age_group, "gender": payload.gender, "roles": payload.roles, "participant_id": payload.participant_id, "message_index": payload.message_index, "rating": payload.rating, "reply_content": payload.reply_content, "reply_id": str(payload.reply_id), }, ) # Endpoint for specific generic comments @app.post("/comment") @limiter.limit("450/minute") def comment_endpoint( payload: CommentRequest, background_tasks: BackgroundTasks, request: Request ): logger.info("Received comment") background_tasks.add_task( log_chat_event, user_id=payload.user_id, session_id=payload.session_id, data={ "consent": payload.consent, "comment": payload.comment, "age_group": payload.age_group, "gender": payload.gender, "roles": payload.roles, "participant_id": payload.participant_id, }, ) @app.put("/file") @limiter.limit("12/minute") async def upload_file( request: Request, file: UploadFile = File(...), session_id: str = Form( pattern="^[a-zA-Z0-9_-]+$", min_length=1, max_length=MAX_ID_LENGTH ), ): try: validated_file = await validate_file(file) except FileValidationException as e: status_code = FILE_VALIDATION_ERROR_STATUS_CODES[e.error] return Response(status_code=status_code) file_content = validated_file.content file_name = validated_file.filename file_mime = validated_file.mime_type try: file_text = await extract_text_from_file(file_content, file_mime) except FileExtractionException as e: status_code = FILE_EXTRACTION_ERROR_STATUS_CODES[e.error] return Response(status_code=status_code) except Exception: # TODO: Log the unexpected failure return Response(status_code=STATUS_CODE_INTERNAL_SERVER_ERROR) try: session_document_store.create_document(session_id, file_text, file_name) except FileDocumentStoreException as e: status_code = FILE_DOCUMENT_STORE_ERROR_STATUS_CODES[e.error] return Response(status_code=status_code) session_tracker.update_session(session_id) @app.delete("/file") @limiter.limit("20/minute") def delete_file( payload: DeleteFileRequest, request: Request, ): session_id = payload.session_id file_name = payload.file_name file_name = replace_spaces_in_filename(file_name) session_document_store.delete_document(session_id, file_name) @app.post("/flush-environmental-infra-impact") @limiter.limit("2/minute") def flush_environmental_infra_impact(request: Request): pass # log_environment_infra()