File size: 5,542 Bytes
a07fdc6 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 | from __future__ import annotations
import logging
import os
from datetime import datetime, timezone
from threading import RLock
from fastapi import Depends, FastAPI, Header, HTTPException, Request, status
from evd_agent.conversation import ConversationManager
from evd_agent.models import InterviewState, InterviewStatus
from .schemas import (
ClassificationResponse,
FinalizeSessionResponse,
HealthResponse,
StartSessionResponse,
SubmitTurnRequest,
SubmitTurnResponse,
)
logger = logging.getLogger("evd_api")
logging.basicConfig(level=os.getenv("EVD_API_LOG_LEVEL", "INFO"))
app = FastAPI(
title="EVD Screening API",
version="1.0.0",
description="Production API surface for EMR integration.",
)
manager = ConversationManager(context_path=os.getenv("EVD_CONTEXT_PATH"))
# In-memory session store for API state. Replace with durable storage in production HA setups.
_session_store: dict[str, InterviewState] = {}
_store_lock = RLock()
@app.middleware("http")
async def audit_log_requests(request: Request, call_next):
# Request-level audit trail with latency and caller metadata.
started = datetime.now(timezone.utc)
response = await call_next(request)
elapsed_ms = int((datetime.now(timezone.utc) - started).total_seconds() * 1000)
logger.info(
"audit request method=%s path=%s status=%s elapsed_ms=%s client=%s",
request.method,
request.url.path,
response.status_code,
elapsed_ms,
request.client.host if request.client else "unknown",
)
return response
def _ensure_authorized(x_api_key: str | None = Header(default=None)) -> None:
# Auth is opt-in: only enforced when EVD_API_KEY is configured.
configured_key = os.getenv("EVD_API_KEY")
if not configured_key:
return
if x_api_key != configured_key:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Unauthorized")
def _get_state_or_404(session_id: str) -> InterviewState:
# Centralized state lookup to keep endpoint handlers simple and consistent.
with _store_lock:
state = _session_store.get(session_id)
if state is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Session not found")
return state
@app.get("/health", response_model=HealthResponse, tags=["health"])
def health() -> HealthResponse:
return HealthResponse(status="ok", service="evd-screening-api")
@app.post("/api/v1/sessions/start", response_model=StartSessionResponse, tags=["sessions"])
def start_session(_: None = Depends(_ensure_authorized)) -> StartSessionResponse:
state, result = manager.start_session()
now = datetime.now(timezone.utc)
with _store_lock:
_session_store[state.session_id] = state
logger.info("audit event=session_started session_id=%s", state.session_id)
return StartSessionResponse(
session_id=state.session_id,
status=state.status.value,
assistant_message=result.assistant_message,
decision=result.decision.model_dump(mode="json"),
risk_profile=result.risk_profile.model_dump(mode="json"),
llm_summary=result.llm_summary,
created_at=now,
)
@app.post("/api/v1/sessions/{session_id}/turn", response_model=SubmitTurnResponse, tags=["sessions"])
def submit_turn(
session_id: str,
payload: SubmitTurnRequest,
_: None = Depends(_ensure_authorized),
) -> SubmitTurnResponse:
state = _get_state_or_404(session_id)
if state.status == InterviewStatus.COMPLETE:
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Session is already finalized")
result = manager.process_turn(state, payload.clinician_message.strip())
logger.info("audit event=turn_submitted session_id=%s", session_id)
return SubmitTurnResponse(
session_id=session_id,
status=state.status.value,
assistant_message=result.assistant_message,
decision=result.decision.model_dump(mode="json"),
risk_profile=result.risk_profile.model_dump(mode="json"),
llm_summary=result.llm_summary,
state_updates=result.state_updates,
)
@app.get(
"/api/v1/sessions/{session_id}/classification",
response_model=ClassificationResponse,
tags=["sessions"],
)
def get_current_classification(
session_id: str,
_: None = Depends(_ensure_authorized),
) -> ClassificationResponse:
state = _get_state_or_404(session_id)
return ClassificationResponse(
session_id=session_id,
status=state.status.value,
decision=state.decision.model_dump(mode="json"),
risk_profile=state.risk_profile.model_dump(mode="json"),
llm_summary=state.llm_summary,
updated_at=datetime.now(timezone.utc),
)
@app.post(
"/api/v1/sessions/{session_id}/finalize",
response_model=FinalizeSessionResponse,
tags=["sessions"],
)
def finalize_session(
session_id: str,
_: None = Depends(_ensure_authorized),
) -> FinalizeSessionResponse:
state = _get_state_or_404(session_id)
state.status = InterviewStatus.COMPLETE
transcript = [turn.model_dump(mode="json") for turn in state.history]
logger.info("audit event=session_finalized session_id=%s", session_id)
return FinalizeSessionResponse(
session_id=session_id,
status=state.status.value,
finalized=True,
final_decision=state.decision.model_dump(mode="json"),
transcript=transcript,
finalized_at=datetime.now(timezone.utc),
)
|