main
py 472 lines 17.6 KB
Raw
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)