File size: 5,103 Bytes
b76f199
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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)