"""Survey product tier (1/2/3) classification from local exemplar corpus + user apply.""" import logging from pathlib import Path from fastapi import APIRouter, Depends, HTTPException, Request from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.api.rate_limit import check_read from app.db.database import get_db from app.db.models import Document from app.models.schemas import ( DocumentSurveyLevelResponse, SurveyLevelApplyRequest, SurveyLevelApplyResponse, SurveyLevelClassifyRequest, SurveyLevelClassifyResponse, SurveyLevelClassificationItem, SurveyLevelCorpusRefreshRequest, SurveyLevelCorpusRefreshResponse, SurveyLevelScoreCandidate, ) from app.services.survey_level_classifier import ( build_or_load_corpus_profiles, classify_file_path, ) logger = logging.getLogger(__name__) router = APIRouter() @router.post( "/documents/survey-level/corpus-refresh", response_model=SurveyLevelCorpusRefreshResponse, summary="Rebuild cached corpus profiles used for tier classification", ) async def corpus_refresh( body: SurveyLevelCorpusRefreshRequest, _: None = Depends(check_read), ) -> SurveyLevelCorpusRefreshResponse: """Scan ``knowledge_base_dirs`` and refresh ``survey_level_corpus_profiles.json``.""" try: prof = build_or_load_corpus_profiles(force_refresh=body.force) return SurveyLevelCorpusRefreshResponse( files_used=prof.files_used, created_at_unix=prof.created_at_unix, ) except Exception as exc: # noqa: BLE001 logger.exception("corpus-refresh failed: %s", exc) raise HTTPException(status_code=500, detail="Corpus refresh failed") from exc @router.post( "/documents/survey-level/classify", response_model=SurveyLevelClassifyResponse, summary="Suggest RICS tier (1/2/3) for uploads using local exemplar corpus", ) async def classify_documents( request: Request, body: SurveyLevelClassifyRequest, db: AsyncSession = Depends(get_db), _: None = Depends(check_read), ) -> SurveyLevelClassifyResponse: """Return predictions only — callers must confirm via ``/documents/survey-level/apply``.""" tenant_id: str = request.state.tenant_id try: profiles = build_or_load_corpus_profiles(force_refresh=body.force_refresh_corpus) except Exception as exc: # noqa: BLE001 logger.exception("classify: corpus build failed: %s", exc) raise HTTPException(status_code=500, detail="Corpus profile build failed") from exc res = await db.execute(select(Document).where(Document.id.in_(body.document_ids))) rows = {d.id: d for d in res.scalars().all()} items: list[SurveyLevelClassificationItem] = [] for doc_id in body.document_ids: doc = rows.get(doc_id) if doc is None or doc.tenant_id != tenant_id: raise HTTPException(status_code=404, detail=f"Document not found: {doc_id}") path = Path(doc.file_path) if not path.is_file(): raise HTTPException( status_code=400, detail=f"Document file missing on disk for {doc_id}; cannot classify.", ) c = classify_file_path(path, filename=doc.filename or path.name, document_id=doc_id, profiles=profiles) items.append( SurveyLevelClassificationItem( document_id=doc_id, filename=c.filename, predicted_survey_level=c.predicted_survey_level, confidence=c.confidence, candidates=[SurveyLevelScoreCandidate(survey_level=x.survey_level, score=x.score) for x in c.candidates], rationale=c.rationale, ) ) n_labelled = sum(int(profiles.levels.get(L, {}).get("n_docs") or 0) for L in (1, 2, 3)) return SurveyLevelClassifyResponse(items=items, corpus_labelled_files=n_labelled) @router.post( "/documents/survey-level/apply", response_model=SurveyLevelApplyResponse, summary="Persist user-confirmed RICS tiers for one or more documents", ) async def apply_survey_levels( request: Request, body: SurveyLevelApplyRequest, db: AsyncSession = Depends(get_db), _: None = Depends(check_read), ) -> SurveyLevelApplyResponse: """Batch PATCH of ``survey_level``; does not re-ingest (retrieval uses DB column).""" tenant_id: str = request.state.tenant_id updated: list[DocumentSurveyLevelResponse] = [] for row in body.items: doc = await db.get(Document, row.document_id) if doc is None or doc.tenant_id != tenant_id: raise HTTPException(status_code=404, detail=f"Document not found: {row.document_id}") doc.survey_level = int(row.survey_level) updated.append( DocumentSurveyLevelResponse( document_id=doc.id, survey_level=doc.survey_level, detail="survey_level saved; future RAG uses this tier for library filtering.", ) ) await db.commit() return SurveyLevelApplyResponse(updated=updated)