37 lines
1.2 KiB
Python
37 lines
1.2 KiB
Python
from app.llm.prompts import CLASSIFIER_PROMPT
|
|
from app.llm.schemas import Classification
|
|
|
|
|
|
def _confidence(value: object) -> float:
|
|
named_levels = {"high": 0.9, "medium": 0.6, "low": 0.3}
|
|
if isinstance(value, str) and value.lower() in named_levels:
|
|
return named_levels[value.lower()]
|
|
try:
|
|
return min(1.0, max(0.0, float(value)))
|
|
except (TypeError, ValueError):
|
|
return 0.0
|
|
|
|
|
|
def _flag(value: object) -> bool:
|
|
if isinstance(value, str):
|
|
return value.strip().lower() in {"true", "yes", "1"}
|
|
return bool(value)
|
|
|
|
|
|
class MessageClassifier:
|
|
def __init__(self, llm):
|
|
self.llm = llm
|
|
|
|
async def classify(self, text: str) -> Classification:
|
|
data = await self.llm.json(
|
|
[{"role": "user", "content": f"{CLASSIFIER_PROMPT}\nMessage:\n{text}"}]
|
|
)
|
|
return Classification(
|
|
is_question=_flag(data.get("is_question")),
|
|
project_related=_flag(data.get("project_related")),
|
|
asks_new_decision=_flag(data.get("asks_new_decision")),
|
|
rhetorical=_flag(data.get("rhetorical")),
|
|
blocker=_flag(data.get("blocker")),
|
|
confidence=_confidence(data.get("confidence", 0)),
|
|
)
|