RICS / app /api /survey_level.py
StormShadow308's picture
- Added support for RICS survey levels (1, 2, 3) in document uploads and reports, allowing for better tier management and retrieval filtering.
b76f199
Raw
History Blame Contribute Delete
5.1 kB
"""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)