131 lines
5.2 KiB
Python
131 lines
5.2 KiB
Python
"""Categorização automática de notas fiscais por fornecedor.
|
|
|
|
Ordem de resolução, sempre retornando um `categoria.id` válido (nunca `None`):
|
|
cache em memória -> limite de chamadas de IA por lote -> IA -> palavra-chave
|
|
-> categoria reservada `1` ("Não Encontrado").
|
|
|
|
Sem retry: qualquer falha na chamada de IA é tratada como "sem categoria da IA"
|
|
e o fluxo cai para o próximo passo, sem nunca impedir a criação do documento.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import sqlite3
|
|
|
|
from . import database as db
|
|
from .config import get_openai_client, get_settings
|
|
|
|
logger = logging.getLogger("lernotafiscal.categorization")
|
|
|
|
# Cache em memória (por fornecedor normalizado) e contador de chamadas de IA
|
|
# por lote — ambos módulo-level, aceitável pois degradam de forma segura
|
|
# (pior caso: recategoriza após restart, ainda limitado pelo teto por lote).
|
|
_supplier_cache: dict[str, int] = {}
|
|
_batch_ai_calls: dict[int, int] = {}
|
|
|
|
|
|
def reset_cache() -> None:
|
|
"""Limpa cache e contadores em memória. Uso principal: isolamento em testes."""
|
|
_supplier_cache.clear()
|
|
_batch_ai_calls.clear()
|
|
|
|
|
|
def _normalize(supplier_name: str) -> str:
|
|
return (supplier_name or "").strip().upper()
|
|
|
|
|
|
def _call_ai(supplier_name: str, known_categories: list[str]) -> str | None:
|
|
"""Chamada única (sem retry) à IA para sugerir uma categoria dentre as
|
|
conhecidas. Qualquer erro ou valor fora do conjunto conhecido -> None."""
|
|
settings = get_settings()
|
|
if not settings.openai_enabled or not known_categories:
|
|
return None
|
|
try:
|
|
client = get_openai_client()
|
|
prompt = (
|
|
"Classifique o fornecedor de uma nota fiscal brasileira em UMA das "
|
|
"categorias a seguir (responda exatamente como escrito na lista): "
|
|
+ ", ".join(known_categories)
|
|
+ f'\nFornecedor: "{supplier_name}"\n'
|
|
'Responda em JSON: {"categoria": "<categoria da lista>"} '
|
|
'ou {"categoria": null} se nenhuma categoria se aplicar claramente.'
|
|
)
|
|
response = client.chat.completions.create(
|
|
model=settings.openai_model,
|
|
messages=[
|
|
{
|
|
"role": "system",
|
|
"content": "Você classifica fornecedores de notas fiscais brasileiras por categoria. Responda SEMPRE em JSON válido.",
|
|
},
|
|
{"role": "user", "content": prompt},
|
|
],
|
|
response_format={"type": "json_object"},
|
|
temperature=0,
|
|
)
|
|
content = response.choices[0].message.content or "{}"
|
|
data = json.loads(content)
|
|
category = data.get("categoria") if isinstance(data, dict) else None
|
|
if isinstance(category, str) and category.strip() in known_categories:
|
|
return category.strip()
|
|
return None
|
|
except Exception as exc: # rede, cota, parsing, modelo indisponível...
|
|
logger.warning("Categorização por IA falhou para '%s': %s", supplier_name, exc)
|
|
return None
|
|
|
|
|
|
def categorize_supplier(supplier_name: str, conn: sqlite3.Connection, batch_id: int) -> int:
|
|
"""Determina o `categoria.id` de um fornecedor. Nunca levanta exceção para o
|
|
chamador nem retorna `None` — em último caso devolve a categoria reservada."""
|
|
try:
|
|
normalized = _normalize(supplier_name)
|
|
|
|
cached = _supplier_cache.get(normalized)
|
|
if cached is not None:
|
|
logger.debug(
|
|
"Categorização de '%s' pulada (cache hit) -> categoria_id=%s", supplier_name, cached
|
|
)
|
|
return cached
|
|
|
|
settings = get_settings()
|
|
categorias = db.list_categorias(conn)
|
|
known_names: list[str] = []
|
|
id_by_name: dict[str, int] = {}
|
|
for row in categorias:
|
|
if row["id"] == db.RESERVED_CATEGORIA_ID:
|
|
continue
|
|
known_names.append(row["categoria"])
|
|
id_by_name.setdefault(row["categoria"], int(row["id"]))
|
|
|
|
categoria_id: int | None = None
|
|
calls_so_far = _batch_ai_calls.get(batch_id, 0)
|
|
limit = settings.categorization_max_ai_calls_per_batch
|
|
|
|
if not settings.openai_enabled:
|
|
pass # sem IA configurada: cai direto para palavra-chave/reservado
|
|
elif calls_so_far >= limit:
|
|
logger.info(
|
|
"Categorização de '%s' pulada (limite de %s chamadas de IA por lote atingido no lote %s)",
|
|
supplier_name, limit, batch_id,
|
|
)
|
|
else:
|
|
_batch_ai_calls[batch_id] = calls_so_far + 1
|
|
ai_category = _call_ai(supplier_name, known_names)
|
|
if ai_category is not None:
|
|
categoria_id = id_by_name.get(ai_category)
|
|
|
|
if categoria_id is None:
|
|
match = db.find_categoria_by_keyword(conn, supplier_name)
|
|
if match is not None:
|
|
categoria_id = int(match["id"])
|
|
|
|
if categoria_id is None:
|
|
categoria_id = db.RESERVED_CATEGORIA_ID
|
|
|
|
_supplier_cache[normalized] = categoria_id
|
|
return categoria_id
|
|
except Exception as exc: # nunca impede a criação do fiscal_documents
|
|
logger.warning("Categorização falhou inesperadamente para '%s': %s", supplier_name, exc)
|
|
return db.RESERVED_CATEGORIA_ID
|