igerasimov's picture
Fix Hugging Face runtime compatibility with Python 3.10
dc54127
Raw
History Blame Contribute Delete
13.2 kB
"""Single-article MVP classification orchestration."""
from __future__ import annotations
import logging
import time
from datetime import datetime, timezone
from typing import Any
from gcmd_classifier.classification import (
ClassificationCandidate,
TermBranchSeed,
candidate_from_terminal_outcome,
candidate_from_topic_stop,
remove_redundant_classifications,
route_terms,
route_topics,
validate_candidates,
)
from gcmd_classifier.config import ModelSettings
from gcmd_classifier.llm.base import ModelClient
from gcmd_classifier.logging_config import get_logger, log_event
from gcmd_classifier.models import (
ArticleClassificationOutcome,
ArticleProcessingStatus,
ArticleRecord,
ArticleResult,
ClassificationRecord,
OutputError,
OutputWarning,
ProcessingMetadata,
ReviewStatus,
)
from gcmd_classifier.persistence.cache import article_fingerprint, configuration_hash
from gcmd_classifier.pipeline.review import flag_review_risk
from gcmd_classifier.vocabulary.index import VocabularyIndex
APPLICATION_VERSION = "0.1.0"
def classify_article(
*,
article: ArticleRecord,
vocabulary: VocabularyIndex,
model_client: ModelClient,
settings: ModelSettings,
run_id: str | None = None,
application_version: str = APPLICATION_VERSION,
relevant_config: dict[str, Any] | None = None,
logger: logging.Logger | None = None,
) -> ArticleResult:
"""Classify one valid article through the current MVP pipeline."""
active_logger = logger or get_logger(__name__)
started_at = _now_iso()
started = time.perf_counter()
model_calls_before = _model_request_count(model_client)
log_event(active_logger, "article_started", DOI=article.DOI, stage="pipeline")
candidates: list[ClassificationCandidate] = []
warnings: list[OutputWarning] = []
errors: list[OutputError] = []
no_classification_reason: str | None = None
try:
topic_result = route_topics(
article=article,
vocabulary=vocabulary,
model_client=model_client,
settings=settings,
)
if topic_result.is_no_topic:
no_classification_reason = (
topic_result.no_selection_reason or "No GCMD Topic was supported by the article."
)
for topic_branch in topic_result.branches:
try:
term_result = route_terms(
article=article,
topic_branch=topic_branch,
vocabulary=vocabulary,
model_client=model_client,
settings=settings,
)
except Exception as exc:
errors.append(_error_from_exception(exc, stage="term_routing", DOI=article.DOI))
log_event(
active_logger,
"branch_failed",
DOI=article.DOI,
stage="term_routing",
branch_id=topic_branch.branch_id,
error_code=exc.__class__.__name__,
)
continue
if term_result.stop_at_topic is not None:
candidates.append(candidate_from_topic_stop(term_result.stop_at_topic))
for term_branch in term_result.term_branches:
_process_term_branch(
article=article,
term_branch=term_branch,
vocabulary=vocabulary,
model_client=model_client,
settings=settings,
candidates=candidates,
warnings=warnings,
errors=errors,
logger=active_logger,
)
except Exception as exc:
errors.append(_error_from_exception(exc, stage="topic_routing", DOI=article.DOI))
log_event(
active_logger,
"article_failed",
DOI=article.DOI,
stage="topic_routing",
error_code=exc.__class__.__name__,
)
validation_result = validate_candidates(candidates, vocabulary)
for rejected in validation_result.rejected:
errors.extend(rejected.deterministic_validation.errors)
redundancy_result = remove_redundant_classifications(validation_result.accepted, vocabulary)
warnings.extend(redundancy_result.warnings)
classifications = flag_review_risk(redundancy_result.classifications)
result = _article_result(
article=article,
classifications=classifications,
warnings=tuple(warnings),
errors=tuple(errors),
no_classification_reason=no_classification_reason,
started_at=started_at,
duration_seconds=time.perf_counter() - started,
model_calls=_model_calls_used(model_client, model_calls_before),
run_id=run_id,
vocabulary=vocabulary,
settings=settings,
application_version=application_version,
relevant_config=relevant_config,
)
log_event(
active_logger,
"article_finished",
DOI=article.DOI,
stage="pipeline",
processing_status=result.processing_status.value,
classification_outcome=None
if result.classification_outcome is None
else result.classification_outcome.value,
classifications=len(result.classifications),
errors=len(result.errors),
)
return result
def _process_term_branch(
*,
article: ArticleRecord,
term_branch: TermBranchSeed,
vocabulary: VocabularyIndex,
model_client: ModelClient,
settings: ModelSettings,
candidates: list[ClassificationCandidate],
warnings: list[OutputWarning],
errors: list[OutputError],
logger: logging.Logger,
) -> None:
from gcmd_classifier.classification import descend_variables
try:
descent_result = descend_variables(
article=article,
term_branch=term_branch,
vocabulary=vocabulary,
model_client=model_client,
settings=settings,
)
except Exception as exc:
errors.append(_error_from_exception(exc, stage="variable_descent", DOI=article.DOI))
log_event(
logger,
"branch_failed",
DOI=article.DOI,
stage="variable_descent",
branch_id=term_branch.branch_id,
error_code=exc.__class__.__name__,
)
return
candidates.extend(
candidate_from_terminal_outcome(terminal) for terminal in descent_result.terminals
)
warnings.extend(descent_result.warnings)
errors.extend(descent_result.diagnostics)
for branch_error in descent_result.errors:
errors.append(
OutputError(
code=branch_error.code,
message=branch_error.message,
stage="variable_descent",
DOI=article.DOI,
details={
"branch_id": branch_error.branch_id,
"parent_branch_id": branch_error.parent_branch_id,
"parent_uuid": branch_error.parent_uuid,
"parent_level": branch_error.parent_level,
},
)
)
def _article_result(
*,
article: ArticleRecord,
classifications: tuple[ClassificationRecord, ...],
warnings: tuple[OutputWarning, ...],
errors: tuple[OutputError, ...],
no_classification_reason: str | None,
started_at: str,
duration_seconds: float,
model_calls: int | None,
run_id: str | None,
vocabulary: VocabularyIndex,
settings: ModelSettings,
application_version: str,
relevant_config: dict[str, Any] | None,
) -> ArticleResult:
metadata = _metadata(
article=article,
started_at=started_at,
duration_seconds=duration_seconds,
model_calls=model_calls,
run_id=run_id,
vocabulary=vocabulary,
settings=settings,
application_version=application_version,
relevant_config=relevant_config,
)
if classifications:
status = ArticleProcessingStatus.PARTIAL if errors else ArticleProcessingStatus.COMPLETED
return ArticleResult(
DOI=article.DOI,
Title=article.Title,
Year=article.Year,
Abstract=article.Abstract,
processing_status=status,
classification_outcome=ArticleClassificationOutcome.CLASSIFIED,
classifications=classifications,
review_status=ReviewStatus.NOT_REQUIRED,
warnings=warnings,
errors=errors,
processing_metadata=metadata,
)
if errors:
return ArticleResult(
DOI=article.DOI,
Title=article.Title,
Year=article.Year,
Abstract=article.Abstract,
processing_status=ArticleProcessingStatus.FAILED,
classification_outcome=None,
classifications=(),
review_status=ReviewStatus.NOT_REQUIRED,
warnings=warnings,
errors=errors,
processing_metadata=metadata,
)
return ArticleResult(
DOI=article.DOI,
Title=article.Title,
Year=article.Year,
Abstract=article.Abstract,
processing_status=ArticleProcessingStatus.COMPLETED,
classification_outcome=ArticleClassificationOutcome.NOT_CLASSIFIED,
classifications=(),
no_classification_reason=no_classification_reason
or "No defensible GCMD classification was supported.",
review_status=ReviewStatus.NOT_REQUIRED,
warnings=warnings,
errors=errors,
processing_metadata=metadata,
)
def failed_article_result(
*,
article: ArticleRecord,
error: OutputError,
vocabulary: VocabularyIndex,
settings: ModelSettings,
run_id: str | None = None,
application_version: str = APPLICATION_VERSION,
relevant_config: dict[str, Any] | None = None,
) -> ArticleResult:
"""Build and return a persisted failed article result for batch-level failures."""
started_at = _now_iso()
return ArticleResult(
DOI=article.DOI,
Title=article.Title,
Year=article.Year,
Abstract=article.Abstract,
processing_status=ArticleProcessingStatus.FAILED,
classification_outcome=None,
classifications=(),
review_status=ReviewStatus.NOT_REQUIRED,
errors=(error,),
processing_metadata=_metadata(
article=article,
started_at=started_at,
duration_seconds=0.0,
model_calls=None,
run_id=run_id,
vocabulary=vocabulary,
settings=settings,
application_version=application_version,
relevant_config=relevant_config,
),
)
def _metadata(
*,
article: ArticleRecord,
started_at: str,
duration_seconds: float,
model_calls: int | None,
run_id: str | None,
vocabulary: VocabularyIndex,
settings: ModelSettings,
application_version: str,
relevant_config: dict[str, Any] | None,
) -> ProcessingMetadata:
completed_at = _now_iso()
return ProcessingMetadata(
run_id=run_id,
started_at=started_at,
completed_at=completed_at,
processed_at=completed_at,
model_provider=settings.provider,
model_name=settings.model_name,
model_parameters={
"temperature": settings.temperature,
"timeout_seconds": settings.timeout_seconds,
"max_retries": settings.max_retries,
},
prompt_versions={
"topic": settings.prompt_version_topic,
"term": settings.prompt_version_term,
"variable": settings.prompt_version_variable,
},
vocabulary_version=vocabulary.vocabulary_version,
vocabulary_hash=vocabulary.vocabulary_version,
application_version=application_version,
configuration_hash=configuration_hash(relevant_config),
article_fingerprint=article_fingerprint(article),
cache_used=False,
processing_time_seconds=duration_seconds,
model_calls=model_calls,
title_available=bool(article.Title),
abstract_available=bool(article.Abstract),
)
def _error_from_exception(exc: Exception, *, stage: str, DOI: str | None = None) -> OutputError:
retry_count = getattr(exc, "retry_count", None)
return OutputError(
code=exc.__class__.__name__,
message=str(exc),
stage=stage,
DOI=DOI,
retry_count=retry_count if isinstance(retry_count, int) else None,
)
def _model_request_count(model_client: ModelClient) -> int | None:
requests = getattr(model_client, "requests", None)
return len(requests) if isinstance(requests, list) else None
def _model_calls_used(model_client: ModelClient, before: int | None) -> int | None:
after = _model_request_count(model_client)
if before is None or after is None:
return None
return after - before
def _now_iso() -> str:
return datetime.now(timezone.utc).isoformat() # noqa: UP017