| 1 | from typing import Optional |
| 2 | |
| 3 | from fastapi import APIRouter |
| 4 | from fastapi import Depends |
| 5 | from fastapi import HTTPException |
| 6 | from fastapi import Query |
| 7 | from fastapi import Security |
| 8 | from loguru import logger |
| 9 | from sqlalchemy.ext.asyncio import AsyncSession |
| 10 | |
| 11 | from app.ai_analyst.schema.ai_analyst import AlertAnalysisResponse |
| 12 | from app.ai_analyst.schema.ai_analyst import AlertsWithReportsListResponse |
| 13 | from app.ai_analyst.schema.ai_analyst import CreateJobRequest |
| 14 | from app.ai_analyst.schema.ai_analyst import CreateJobResponse |
| 15 | from app.ai_analyst.schema.ai_analyst import IocListResponse |
| 16 | from app.ai_analyst.schema.ai_analyst import JobListResponse |
| 17 | from app.ai_analyst.schema.ai_analyst import MyReviewResponse |
| 18 | from app.ai_analyst.schema.ai_analyst import PalaceConsolidationResponse |
| 19 | from app.ai_analyst.schema.ai_analyst import PalaceSearchResponse |
| 20 | from app.ai_analyst.schema.ai_analyst import QueuePalaceLessonRequest |
| 21 | from app.ai_analyst.schema.ai_analyst import QueuePalaceLessonResponse |
| 22 | from app.ai_analyst.schema.ai_analyst import ReplayRequest |
| 23 | from app.ai_analyst.schema.ai_analyst import ReplayResponse |
| 24 | from app.ai_analyst.schema.ai_analyst import ReportListResponse |
| 25 | from app.ai_analyst.schema.ai_analyst import ReviewListResponse |
| 26 | from app.ai_analyst.schema.ai_analyst import ReviewStatsResponse |
| 27 | from app.ai_analyst.schema.ai_analyst import SubmitIocsRequest |
| 28 | from app.ai_analyst.schema.ai_analyst import SubmitIocsResponse |
| 29 | from app.ai_analyst.schema.ai_analyst import SubmitReportRequest |
| 30 | from app.ai_analyst.schema.ai_analyst import SubmitReportResponse |
| 31 | from app.ai_analyst.schema.ai_analyst import SubmitReviewRequest |
| 32 | from app.ai_analyst.schema.ai_analyst import SubmitReviewResponse |
| 33 | from app.ai_analyst.schema.ai_analyst import UpdateJobRequest |
| 34 | from app.ai_analyst.schema.ai_analyst import UpdateJobResponse |
| 35 | from app.ai_analyst.services.ai_analyst import create_job |
| 36 | from app.ai_analyst.services.ai_analyst import get_alert_analysis |
| 37 | from app.ai_analyst.services.ai_analyst import get_job |
| 38 | from app.ai_analyst.services.ai_analyst import get_my_review |
| 39 | from app.ai_analyst.services.ai_analyst import get_palace_consolidation |
| 40 | from app.ai_analyst.services.ai_analyst import get_review_stats |
| 41 | from app.ai_analyst.services.ai_analyst import list_alerts_with_reports |
| 42 | from app.ai_analyst.services.ai_analyst import list_iocs_by_alert |
| 43 | from app.ai_analyst.services.ai_analyst import list_iocs_by_customer |
| 44 | from app.ai_analyst.services.ai_analyst import list_iocs_by_report |
| 45 | from app.ai_analyst.services.ai_analyst import list_jobs_by_alert |
| 46 | from app.ai_analyst.services.ai_analyst import list_jobs_by_customer |
| 47 | from app.ai_analyst.services.ai_analyst import list_reports_by_alert |
| 48 | from app.ai_analyst.services.ai_analyst import list_reviews_by_customer |
| 49 | from app.ai_analyst.services.ai_analyst import queue_palace_lesson |
| 50 | from app.ai_analyst.services.ai_analyst import submit_iocs |
| 51 | from app.ai_analyst.services.ai_analyst import submit_report |
| 52 | from app.ai_analyst.services.ai_analyst import submit_review |
| 53 | from app.ai_analyst.services.ai_analyst import update_job |
| 54 | from app.auth.models.users import User |
| 55 | from app.auth.utils import AuthHandler |
| 56 | from app.connectors.talon.services.talon import ( |
| 57 | replay_investigation as talon_replay_investigation, |
| 58 | ) |
| 59 | from app.connectors.talon.services.talon import ( |
| 60 | search_palace_lessons as talon_search_palace_lessons, |
| 61 | ) |
| 62 | from app.db.db_session import get_db |
| 63 | from app.db.universal_models import AiAnalystReport |
| 64 | |
| 65 | ai_analyst_router = APIRouter() |
| 66 | |
| 67 | |
| 68 | # --- Job endpoints --- |
| 69 | |
| 70 | |
| 71 | @ai_analyst_router.post( |
| 72 | "/jobs", |
| 73 | response_model=CreateJobResponse, |
| 74 | description="Register a new AI analyst investigation job", |
| 75 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 76 | ) |
| 77 | async def create_job_route( |
| 78 | request: CreateJobRequest, |
| 79 | session: AsyncSession = Depends(get_db), |
| 80 | ) -> CreateJobResponse: |
| 81 | logger.info(f"Creating AI analyst job for alert {request.alert_id}") |
| 82 | return await create_job(request, session) |
| 83 | |
| 84 | |
| 85 | @ai_analyst_router.patch( |
| 86 | "/jobs/{job_id}", |
| 87 | response_model=UpdateJobResponse, |
| 88 | description="Update an AI analyst job status", |
| 89 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 90 | ) |
| 91 | async def update_job_route( |
| 92 | job_id: str, |
| 93 | request: UpdateJobRequest, |
| 94 | session: AsyncSession = Depends(get_db), |
| 95 | ) -> UpdateJobResponse: |
| 96 | logger.info(f"Updating AI analyst job {job_id}") |
| 97 | return await update_job(job_id, request, session) |
| 98 | |
| 99 | |
| 100 | @ai_analyst_router.get( |
| 101 | "/jobs/{job_id}", |
| 102 | response_model=CreateJobResponse, |
| 103 | description="Get a specific AI analyst job", |
| 104 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 105 | ) |
| 106 | async def get_job_route( |
| 107 | job_id: str, |
| 108 | session: AsyncSession = Depends(get_db), |
| 109 | ) -> CreateJobResponse: |
| 110 | job = await get_job(job_id, session) |
| 111 | return CreateJobResponse(success=True, message="Job retrieved", job=job) |
| 112 | |
| 113 | |
| 114 | @ai_analyst_router.get( |
| 115 | "/jobs/alert/{alert_id}", |
| 116 | response_model=JobListResponse, |
| 117 | description="List all AI analyst jobs for an alert", |
| 118 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 119 | ) |
| 120 | async def list_jobs_by_alert_route( |
| 121 | alert_id: int, |
| 122 | session: AsyncSession = Depends(get_db), |
| 123 | ) -> JobListResponse: |
| 124 | jobs = await list_jobs_by_alert(alert_id, session) |
| 125 | return JobListResponse(success=True, message="Jobs retrieved", jobs=jobs) |
| 126 | |
| 127 | |
| 128 | @ai_analyst_router.get( |
| 129 | "/jobs/customer/{customer_code}", |
| 130 | response_model=JobListResponse, |
| 131 | description="List all AI analyst jobs for a customer", |
| 132 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 133 | ) |
| 134 | async def list_jobs_by_customer_route( |
| 135 | customer_code: str, |
| 136 | session: AsyncSession = Depends(get_db), |
| 137 | ) -> JobListResponse: |
| 138 | jobs = await list_jobs_by_customer(customer_code, session) |
| 139 | return JobListResponse(success=True, message="Jobs retrieved", jobs=jobs) |
| 140 | |
| 141 | |
| 142 | # --- Report endpoints --- |
| 143 | |
| 144 | |
| 145 | @ai_analyst_router.post( |
| 146 | "/reports", |
| 147 | response_model=SubmitReportResponse, |
| 148 | description="Submit an AI analyst investigation report", |
| 149 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 150 | ) |
| 151 | async def submit_report_route( |
| 152 | request: SubmitReportRequest, |
| 153 | session: AsyncSession = Depends(get_db), |
| 154 | ) -> SubmitReportResponse: |
| 155 | logger.info(f"Submitting AI analyst report for job {request.job_id}") |
| 156 | return await submit_report(request, session) |
| 157 | |
| 158 | |
| 159 | @ai_analyst_router.get( |
| 160 | "/reports/alert/{alert_id}", |
| 161 | response_model=ReportListResponse, |
| 162 | description="List all AI analyst reports for an alert", |
| 163 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 164 | ) |
| 165 | async def list_reports_by_alert_route( |
| 166 | alert_id: int, |
| 167 | session: AsyncSession = Depends(get_db), |
| 168 | ) -> ReportListResponse: |
| 169 | reports = await list_reports_by_alert(alert_id, session) |
| 170 | return ReportListResponse(success=True, message="Reports retrieved", reports=reports) |
| 171 | |
| 172 | |
| 173 | # --- IOC endpoints --- |
| 174 | |
| 175 | |
| 176 | @ai_analyst_router.post( |
| 177 | "/iocs", |
| 178 | response_model=SubmitIocsResponse, |
| 179 | description="Submit extracted IOCs for a report", |
| 180 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 181 | ) |
| 182 | async def submit_iocs_route( |
| 183 | request: SubmitIocsRequest, |
| 184 | session: AsyncSession = Depends(get_db), |
| 185 | ) -> SubmitIocsResponse: |
| 186 | logger.info(f"Submitting IOCs for report {request.report_id}") |
| 187 | return await submit_iocs(request, session) |
| 188 | |
| 189 | |
| 190 | @ai_analyst_router.get( |
| 191 | "/iocs/report/{report_id}", |
| 192 | response_model=IocListResponse, |
| 193 | description="List IOCs for a specific report", |
| 194 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 195 | ) |
| 196 | async def list_iocs_by_report_route( |
| 197 | report_id: int, |
| 198 | session: AsyncSession = Depends(get_db), |
| 199 | ) -> IocListResponse: |
| 200 | iocs = await list_iocs_by_report(report_id, session) |
| 201 | return IocListResponse(success=True, message="IOCs retrieved", iocs=iocs) |
| 202 | |
| 203 | |
| 204 | @ai_analyst_router.get( |
| 205 | "/iocs/alert/{alert_id}", |
| 206 | response_model=IocListResponse, |
| 207 | description="List all IOCs for an alert", |
| 208 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 209 | ) |
| 210 | async def list_iocs_by_alert_route( |
| 211 | alert_id: int, |
| 212 | session: AsyncSession = Depends(get_db), |
| 213 | ) -> IocListResponse: |
| 214 | iocs = await list_iocs_by_alert(alert_id, session) |
| 215 | return IocListResponse(success=True, message="IOCs retrieved", iocs=iocs) |
| 216 | |
| 217 | |
| 218 | @ai_analyst_router.get( |
| 219 | "/iocs/customer/{customer_code}", |
| 220 | response_model=IocListResponse, |
| 221 | description="List IOCs for a customer, optionally filtered by VT verdict", |
| 222 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 223 | ) |
| 224 | async def list_iocs_by_customer_route( |
| 225 | customer_code: str, |
| 226 | vt_verdict: Optional[str] = Query(None, description="Filter by VirusTotal verdict"), |
| 227 | session: AsyncSession = Depends(get_db), |
| 228 | ) -> IocListResponse: |
| 229 | iocs = await list_iocs_by_customer(customer_code, session, vt_verdict=vt_verdict) |
| 230 | return IocListResponse(success=True, message="IOCs retrieved", iocs=iocs) |
| 231 | |
| 232 | |
| 233 | # --- Alerts with reports --- |
| 234 | |
| 235 | |
| 236 | @ai_analyst_router.get( |
| 237 | "/alerts_with_reports", |
| 238 | response_model=AlertsWithReportsListResponse, |
| 239 | description="List all alerts that have an AI analyst report", |
| 240 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 241 | ) |
| 242 | async def list_alerts_with_reports_route( |
| 243 | customer_code: Optional[str] = Query(None, description="Filter by customer code"), |
| 244 | session: AsyncSession = Depends(get_db), |
| 245 | ) -> AlertsWithReportsListResponse: |
| 246 | alerts = await list_alerts_with_reports(session, customer_code=customer_code) |
| 247 | return AlertsWithReportsListResponse( |
| 248 | success=True, |
| 249 | message=f"{len(alerts)} alerts with reports found", |
| 250 | alerts=alerts, |
| 251 | ) |
| 252 | |
| 253 | |
| 254 | # --- Combined alert analysis endpoint --- |
| 255 | |
| 256 | |
| 257 | @ai_analyst_router.get( |
| 258 | "/alert/{alert_id}", |
| 259 | response_model=AlertAnalysisResponse, |
| 260 | description="Get the full AI analysis for an alert (job, report, IOCs)", |
| 261 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 262 | ) |
| 263 | async def get_alert_analysis_route( |
| 264 | alert_id: int, |
| 265 | session: AsyncSession = Depends(get_db), |
| 266 | ) -> AlertAnalysisResponse: |
| 267 | job, report, iocs = await get_alert_analysis(alert_id, session) |
| 268 | if not job: |
| 269 | return AlertAnalysisResponse(success=False, message="No AI analysis found for this alert") |
| 270 | return AlertAnalysisResponse( |
| 271 | success=True, |
| 272 | message="Alert analysis retrieved", |
| 273 | job=job, |
| 274 | report=report, |
| 275 | iocs=iocs, |
| 276 | ) |
| 277 | |
| 278 | |
| 279 | # --- Review / Palace lesson / Replay endpoints --- |
| 280 | |
| 281 | |
| 282 | @ai_analyst_router.post( |
| 283 | "/reports/{report_id}/review", |
| 284 | response_model=SubmitReviewResponse, |
| 285 | description="Submit an analyst review (rubric + IOC corrections) for a report", |
| 286 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 287 | ) |
| 288 | async def submit_review_route( |
| 289 | report_id: int, |
| 290 | request: SubmitReviewRequest, |
| 291 | current_user: User = Depends(AuthHandler().get_current_user), |
| 292 | session: AsyncSession = Depends(get_db), |
| 293 | ) -> SubmitReviewResponse: |
| 294 | """ |
| 295 | Persist an analyst review rubric + per-IOC corrections. Scope gated |
| 296 | (admin OR analyst) via require_any_scope; the authenticated user's id is |
| 297 | captured as reviewer_user_id for audit. |
| 298 | """ |
| 299 | logger.info(f"User {current_user.id} submitting review for report {report_id}") |
| 300 | return await submit_review( |
| 301 | report_id=report_id, |
| 302 | request=request, |
| 303 | reviewer_user_id=current_user.id, |
| 304 | session=session, |
| 305 | ) |
| 306 | |
| 307 | |
| 308 | @ai_analyst_router.get( |
| 309 | "/reports/{report_id}/review/mine", |
| 310 | response_model=MyReviewResponse, |
| 311 | description="Fetch the current user's existing review for a report (returns review=null if none yet)", |
| 312 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 313 | ) |
| 314 | async def get_my_review_route( |
| 315 | report_id: int, |
| 316 | current_user: User = Depends(AuthHandler().get_current_user), |
| 317 | session: AsyncSession = Depends(get_db), |
| 318 | ) -> MyReviewResponse: |
| 319 | """UI calls this on open to decide between create-mode and edit-existing-mode.""" |
| 320 | return await get_my_review( |
| 321 | report_id=report_id, |
| 322 | reviewer_user_id=current_user.id, |
| 323 | session=session, |
| 324 | ) |
| 325 | |
| 326 | |
| 327 | @ai_analyst_router.post( |
| 328 | "/reports/{report_id}/replay", |
| 329 | response_model=ReplayResponse, |
| 330 | description="Replay an investigation for the given report's alert with a forced template override", |
| 331 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 332 | ) |
| 333 | async def replay_report_route( |
| 334 | report_id: int, |
| 335 | request: ReplayRequest, |
| 336 | session: AsyncSession = Depends(get_db), |
| 337 | ) -> ReplayResponse: |
| 338 | """ |
| 339 | Triggers Talon POST /investigate with template_override. The new run will |
| 340 | create its own AiAnalystJob/Report via Talon's existing webhook callbacks — |
| 341 | this endpoint does not mutate local DB itself. |
| 342 | """ |
| 343 | report = await session.get(AiAnalystReport, report_id) |
| 344 | if not report: |
| 345 | raise HTTPException(status_code=404, detail=f"Report {report_id} not found") |
| 346 | |
| 347 | # Guard: customer_code from the client must match the report's, preventing |
| 348 | # replay injection across tenants |
| 349 | if request.customer_code != report.customer_code: |
| 350 | raise HTTPException( |
| 351 | status_code=400, |
| 352 | detail=(f"customer_code mismatch: report belongs to {report.customer_code}, " f"got {request.customer_code}"), |
| 353 | ) |
| 354 | |
| 355 | logger.info( |
| 356 | f"Replaying investigation for report {report_id} (alert {report.alert_id}) " f"with template_override={request.template_override}", |
| 357 | ) |
| 358 | talon_response = await talon_replay_investigation( |
| 359 | alert_id=report.alert_id, |
| 360 | customer_code=request.customer_code, |
| 361 | template_override=request.template_override, |
| 362 | sender=request.sender, |
| 363 | ) |
| 364 | return ReplayResponse( |
| 365 | success=True, |
| 366 | message="Replay triggered", |
| 367 | data=talon_response.get("data"), |
| 368 | ) |
| 369 | |
| 370 | |
| 371 | @ai_analyst_router.post( |
| 372 | "/palace_lessons", |
| 373 | response_model=QueuePalaceLessonResponse, |
| 374 | description="Queue a MemPalace lesson for async ingestion by the NanoClaw drainer", |
| 375 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 376 | ) |
| 377 | async def queue_palace_lesson_route( |
| 378 | request: QueuePalaceLessonRequest, |
| 379 | session: AsyncSession = Depends(get_db), |
| 380 | ) -> QueuePalaceLessonResponse: |
| 381 | logger.info(f"Queuing palace lesson for customer {request.customer_code}") |
| 382 | return await queue_palace_lesson(request, session) |
| 383 | |
| 384 | |
| 385 | @ai_analyst_router.get( |
| 386 | "/reviews/customer/{customer_code}", |
| 387 | response_model=ReviewListResponse, |
| 388 | description="Review dashboard feed for a customer (newest first, with per-IOC reviews)", |
| 389 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 390 | ) |
| 391 | async def list_reviews_by_customer_route( |
| 392 | customer_code: str, |
| 393 | session: AsyncSession = Depends(get_db), |
| 394 | ) -> ReviewListResponse: |
| 395 | reviews = await list_reviews_by_customer(customer_code, session) |
| 396 | return ReviewListResponse( |
| 397 | success=True, |
| 398 | message=f"{len(reviews)} reviews retrieved", |
| 399 | reviews=reviews, |
| 400 | ) |
| 401 | |
| 402 | |
| 403 | @ai_analyst_router.get( |
| 404 | "/reviews/customer/{customer_code}/stats", |
| 405 | response_model=ReviewStatsResponse, |
| 406 | description="Aggregate review metrics (feedback dashboard) for a customer — SQL-side rollup", |
| 407 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 408 | ) |
| 409 | async def get_review_stats_route( |
| 410 | customer_code: str, |
| 411 | recent_limit: int = Query(10, ge=0, le=50, description="How many recent reviews to embed"), |
| 412 | session: AsyncSession = Depends(get_db), |
| 413 | ) -> ReviewStatsResponse: |
| 414 | logger.info(f"Fetching review stats for customer {customer_code}") |
| 415 | return await get_review_stats( |
| 416 | customer_code=customer_code, |
| 417 | session=session, |
| 418 | recent_limit=recent_limit, |
| 419 | ) |
| 420 | |
| 421 | |
| 422 | @ai_analyst_router.get( |
| 423 | "/palace_lessons/customer/{customer_code}", |
| 424 | response_model=PalaceSearchResponse, |
| 425 | description="Preview similar MemPalace lessons for a customer (proxies to Talon /palace/search)", |
| 426 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 427 | ) |
| 428 | async def search_palace_lessons_route( |
| 429 | customer_code: str, |
| 430 | query: str = Query(..., min_length=1, description="Semantic search query"), |
| 431 | room: Optional[str] = Query(None, description="Optional room filter"), |
| 432 | limit: int = Query(5, ge=1, le=25, description="Max hits"), |
| 433 | ) -> PalaceSearchResponse: |
| 434 | logger.info( |
| 435 | f"Searching palace for customer={customer_code} query={query!r} room={room} limit={limit}", |
| 436 | ) |
| 437 | talon_response = await talon_search_palace_lessons( |
| 438 | customer_code=customer_code, |
| 439 | query=query, |
| 440 | room=room, |
| 441 | limit=limit, |
| 442 | ) |
| 443 | # Talon returns {data: {...}} — try to extract the lessons list regardless of shape |
| 444 | data = talon_response.get("data") or {} |
| 445 | raw_lessons = data.get("lessons") if isinstance(data, dict) else None |
| 446 | if raw_lessons is None and isinstance(data, list): |
| 447 | raw_lessons = data |
| 448 | if raw_lessons is None: |
| 449 | raw_lessons = [] |
| 450 | return PalaceSearchResponse( |
| 451 | success=True, |
| 452 | message=f"{len(raw_lessons)} palace lessons matched", |
| 453 | lessons=raw_lessons, |
| 454 | ) |
| 455 | |
| 456 | |
| 457 | @ai_analyst_router.get( |
| 458 | "/palace_lessons/customer/{customer_code}/consolidation", |
| 459 | response_model=PalaceConsolidationResponse, |
| 460 | description=( |
| 461 | "Manual consolidation digest for a customer's active MemPalace lessons — " |
| 462 | "groups by room, flags near-duplicate pairs, and surfaces one-offs about " |
| 463 | "to be swept. Pure read-only; no Talon round-trip." |
| 464 | ), |
| 465 | dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst"))], |
| 466 | ) |
| 467 | async def get_palace_consolidation_route( |
| 468 | customer_code: str, |
| 469 | session: AsyncSession = Depends(get_db), |
| 470 | ) -> PalaceConsolidationResponse: |
| 471 | logger.info(f"Building palace consolidation for customer {customer_code}") |
| 472 | return await get_palace_consolidation(customer_code, session) |