| """Controlled Variable-level descent through direct GCMD children.""" |
|
|
| from __future__ import annotations |
|
|
| from pydantic import BaseModel, ConfigDict, Field |
|
|
| from gcmd_classifier.classification.candidates import ( |
| VariableCandidate, |
| build_variable_candidates, |
| validate_variable_candidate_relationship, |
| ) |
| from gcmd_classifier.classification.routers import ModelCallMetadata, TermBranchSeed |
| from gcmd_classifier.config import ModelSettings |
| from gcmd_classifier.errors import UnknownCandidateIDError |
| from gcmd_classifier.llm.base import ( |
| ModelClient, |
| ModelRequest, |
| ModelStage, |
| RetryPolicy, |
| generate_with_retries, |
| ) |
| from gcmd_classifier.llm.prompts import ParentContext, build_variable_prompt |
| from gcmd_classifier.llm.schemas import CandidateDecision, VariableResponse |
| from gcmd_classifier.models import ( |
| ArticleRecord, |
| HierarchyLevel, |
| OutputError, |
| OutputWarning, |
| SupportType, |
| ) |
| from gcmd_classifier.vocabulary.index import VocabularyIndex |
|
|
| _VARIABLE_PARENT_LEVELS = {"Term", "Variable_Level_1", "Variable_Level_2"} |
| _VARIABLE_LEVELS = {"Variable_Level_1", "Variable_Level_2", "Variable_Level_3"} |
|
|
|
|
| class EvidenceStep(BaseModel): |
| """Evidence and routing metadata for one stage in a branch lineage.""" |
|
|
| model_config = ConfigDict(extra="forbid", frozen=True) |
|
|
| branch_id: str = Field(min_length=1) |
| UUID: str = Field(min_length=1) |
| name: str = Field(min_length=1) |
| level: str = Field(min_length=1) |
| evidence: str = Field(min_length=1) |
| support_type: SupportType |
| confidence: float | None = Field(default=None, ge=0.0, le=1.0) |
| reason: str | None = None |
| candidate_id: str | None = None |
|
|
|
|
| class VariableTerminalOutcome(BaseModel): |
| """Terminal outcome for one independently descended branch.""" |
|
|
| model_config = ConfigDict(extra="forbid", frozen=True) |
|
|
| branch_id: str = Field(min_length=1) |
| parent_branch_id: str = Field(min_length=1) |
| topic_uuid: str = Field(min_length=1) |
| topic_name: str = Field(min_length=1) |
| term_uuid: str = Field(min_length=1) |
| term_name: str = Field(min_length=1) |
| final_uuid: str = Field(min_length=1) |
| final_name: str = Field(min_length=1) |
| final_level: HierarchyLevel |
| final_canonical_path: str = Field(min_length=1) |
| path_components: tuple[str, ...] = Field(min_length=1) |
| evidence: str = Field(min_length=1) |
| evidence_trail: tuple[EvidenceStep, ...] = Field(default_factory=tuple) |
| support_type: SupportType |
| confidence: float | None = Field(default=None, ge=0.0, le=1.0) |
| reason: str | None = None |
| stop_reason: str | None = None |
| candidate_id: str | None = None |
| prompt_version: str | None = None |
| model_provider: str | None = None |
| model_name: str | None = None |
| retry_count: int = Field(default=0, ge=0) |
|
|
|
|
| class VariableBranchError(BaseModel): |
| """Explicit branch-level failure that does not cancel sibling branches.""" |
|
|
| model_config = ConfigDict(extra="forbid", frozen=True) |
|
|
| branch_id: str = Field(min_length=1) |
| parent_branch_id: str = Field(min_length=1) |
| parent_uuid: str = Field(min_length=1) |
| parent_name: str = Field(min_length=1) |
| parent_level: str = Field(min_length=1) |
| code: str = Field(min_length=1) |
| message: str = Field(min_length=1) |
|
|
|
|
| class VariableDescentResult(BaseModel): |
| """Result of controlled Variable-level descent for one Term branch.""" |
|
|
| model_config = ConfigDict(extra="forbid", frozen=True) |
|
|
| terminals: tuple[VariableTerminalOutcome, ...] = Field(default_factory=tuple) |
| errors: tuple[VariableBranchError, ...] = Field(default_factory=tuple) |
| warnings: tuple[OutputWarning, ...] = Field(default_factory=tuple) |
| diagnostics: tuple[OutputError, ...] = Field(default_factory=tuple) |
|
|
| @property |
| def terminal_count(self) -> int: |
| """Number of successful terminal branch outcomes.""" |
| return len(self.terminals) |
|
|
| @property |
| def has_errors(self) -> bool: |
| """Whether any independent branch failed during descent.""" |
| return bool(self.errors) |
|
|
|
|
| def descend_variables( |
| *, |
| article: ArticleRecord, |
| term_branch: TermBranchSeed, |
| vocabulary: VocabularyIndex, |
| model_client: ModelClient, |
| settings: ModelSettings, |
| retry_policy: RetryPolicy | None = None, |
| ) -> VariableDescentResult: |
| """Descend from a selected Term branch through supported direct Variable children.""" |
| term = vocabulary.get(term_branch.term_uuid) |
| trail = (_evidence_from_term_branch(term_branch),) |
| context = _BranchContext( |
| topic_uuid=term_branch.parent_topic_uuid, |
| topic_name=term_branch.parent_topic_name, |
| term_uuid=term_branch.term_uuid, |
| term_name=term_branch.term_name, |
| ) |
| terminals, errors = _descend_parent( |
| article=article, |
| parent_uuid=term.UUID, |
| parent_branch_id=term_branch.branch_id, |
| branch_id=term_branch.branch_id, |
| candidate_id=term_branch.candidate_id, |
| parent_evidence=term_branch.evidence, |
| parent_support_type=term_branch.support_type, |
| parent_confidence=term_branch.confidence, |
| parent_reason=term_branch.reason, |
| trail=trail, |
| context=context, |
| vocabulary=vocabulary, |
| model_client=model_client, |
| settings=settings, |
| retry_policy=retry_policy or RetryPolicy.from_settings(settings), |
| parent_metadata=None, |
| ) |
| return VariableDescentResult(terminals=terminals, errors=errors) |
|
|
|
|
| class _BranchContext(BaseModel): |
| model_config = ConfigDict(extra="forbid", frozen=True) |
|
|
| topic_uuid: str |
| topic_name: str |
| term_uuid: str |
| term_name: str |
|
|
|
|
| def _descend_parent( |
| *, |
| article: ArticleRecord, |
| parent_uuid: str, |
| parent_branch_id: str, |
| branch_id: str, |
| candidate_id: str | None, |
| parent_evidence: str, |
| parent_support_type: SupportType, |
| parent_confidence: float | None, |
| parent_reason: str | None, |
| trail: tuple[EvidenceStep, ...], |
| context: _BranchContext, |
| vocabulary: VocabularyIndex, |
| model_client: ModelClient, |
| settings: ModelSettings, |
| retry_policy: RetryPolicy, |
| parent_metadata: ModelCallMetadata | None, |
| ) -> tuple[tuple[VariableTerminalOutcome, ...], tuple[VariableBranchError, ...]]: |
| parent = vocabulary.get(parent_uuid) |
| _validate_variable_parent(parent.level) |
| if parent.level == "Variable_Level_3": |
| return ( |
| ( |
| _terminal_outcome( |
| parent=parent, |
| branch_id=branch_id, |
| parent_branch_id=parent_branch_id, |
| context=context, |
| evidence=parent_evidence, |
| support_type=parent_support_type, |
| confidence=parent_confidence, |
| reason=parent_reason, |
| stop_reason="Selected Variable_Level_3 is a leaf node.", |
| candidate_id=candidate_id, |
| trail=trail, |
| metadata=parent_metadata, |
| ), |
| ), |
| (), |
| ) |
| candidates = build_variable_candidates(vocabulary, parent_uuid=parent_uuid) |
| if not candidates: |
| return ( |
| ( |
| _terminal_outcome( |
| parent=parent, |
| branch_id=branch_id, |
| parent_branch_id=parent_branch_id, |
| context=context, |
| evidence=parent_evidence, |
| support_type=parent_support_type, |
| confidence=parent_confidence, |
| reason=parent_reason, |
| stop_reason="No direct Variable children are available.", |
| candidate_id=candidate_id, |
| trail=trail, |
| metadata=parent_metadata, |
| ), |
| ), |
| (), |
| ) |
|
|
| candidates_by_id = {candidate.candidate_id: candidate for candidate in candidates} |
| response = _call_variable_model( |
| article=article, |
| parent=parent, |
| candidates=candidates, |
| vocabulary=vocabulary, |
| model_client=model_client, |
| settings=settings, |
| retry_policy=retry_policy, |
| ) |
| _validate_selected_candidate_ids(response.parsed.selected, candidates_by_id) |
| metadata = _model_metadata(response) |
| if response.parsed.stop_at_parent: |
| return ( |
| ( |
| _terminal_outcome( |
| parent=parent, |
| branch_id=branch_id, |
| parent_branch_id=parent_branch_id, |
| context=context, |
| evidence=parent_evidence, |
| support_type=parent_support_type, |
| confidence=parent_confidence, |
| reason=parent_reason, |
| stop_reason=response.parsed.stop_reason, |
| candidate_id=candidate_id, |
| trail=trail, |
| metadata=metadata, |
| ), |
| ), |
| (), |
| ) |
|
|
| terminals: list[VariableTerminalOutcome] = [] |
| errors: list[VariableBranchError] = [] |
| for decision in response.parsed.selected: |
| candidate = candidates_by_id[decision.candidate_id] |
| child = vocabulary.get(candidate.variable_uuid) |
| child_branch_id = f"{branch_id}/variable:{candidate.candidate_id}" |
| try: |
| validate_variable_candidate_relationship( |
| candidate, |
| selected_parent_uuid=parent_uuid, |
| index=vocabulary, |
| ) |
| child_trail = ( |
| *trail, |
| _evidence_from_child( |
| child=child, |
| branch_id=child_branch_id, |
| decision=decision, |
| ), |
| ) |
| child_terminals, child_errors = _descend_parent( |
| article=article, |
| parent_uuid=child.UUID, |
| parent_branch_id=branch_id, |
| branch_id=child_branch_id, |
| candidate_id=decision.candidate_id, |
| parent_evidence=decision.evidence, |
| parent_support_type=decision.support_type, |
| parent_confidence=decision.confidence, |
| parent_reason=decision.reason, |
| trail=child_trail, |
| context=context, |
| vocabulary=vocabulary, |
| model_client=model_client, |
| settings=settings, |
| retry_policy=retry_policy, |
| parent_metadata=metadata, |
| ) |
| terminals.extend(child_terminals) |
| errors.extend(child_errors) |
| except Exception as exc: |
| errors.append( |
| _branch_error( |
| branch_id=child_branch_id, |
| parent_branch_id=branch_id, |
| parent=child, |
| exc=exc, |
| ) |
| ) |
| return tuple(terminals), tuple(errors) |
|
|
|
|
| def _call_variable_model( |
| *, |
| article: ArticleRecord, |
| parent, |
| candidates: tuple[VariableCandidate, ...], |
| vocabulary: VocabularyIndex, |
| model_client: ModelClient, |
| settings: ModelSettings, |
| retry_policy: RetryPolicy, |
| ): |
| parent_context = ParentContext( |
| candidate_id=parent.UUID, |
| name=parent.name, |
| level=parent.level, |
| canonical_path=parent.canonical_path, |
| ) |
| prompt = build_variable_prompt( |
| article=article, |
| parent=parent_context, |
| candidates=tuple(candidate.prompt_candidate for candidate in candidates), |
| prompt_version=settings.prompt_version_variable, |
| ) |
| request = ModelRequest.from_settings( |
| stage=ModelStage.VARIABLE, |
| prompt=prompt, |
| response_schema=VariableResponse, |
| settings=settings, |
| metadata={ |
| "DOI": article.DOI, |
| "parent_uuid": parent.UUID, |
| "parent_level": parent.level, |
| "candidate_ids": tuple(candidate.candidate_id for candidate in candidates), |
| "candidate_uuids": tuple(candidate.variable_uuid for candidate in candidates), |
| "vocabulary_version": vocabulary.vocabulary_version, |
| }, |
| ) |
| return generate_with_retries(model_client, request, retry_policy) |
|
|
|
|
| def _validate_selected_candidate_ids( |
| selected: list[CandidateDecision], |
| candidates_by_id: dict[str, VariableCandidate], |
| ) -> None: |
| seen: set[str] = set() |
| for decision in selected: |
| if decision.candidate_id not in candidates_by_id: |
| raise UnknownCandidateIDError( |
| f"Variable model selected unknown candidate_id {decision.candidate_id!r}." |
| ) |
| if decision.candidate_id in seen: |
| raise UnknownCandidateIDError( |
| f"Variable model selected duplicate candidate_id {decision.candidate_id!r}." |
| ) |
| seen.add(decision.candidate_id) |
|
|
|
|
| def _terminal_outcome( |
| *, |
| parent, |
| branch_id: str, |
| parent_branch_id: str, |
| context: _BranchContext, |
| evidence: str, |
| support_type: SupportType, |
| confidence: float | None, |
| reason: str | None, |
| stop_reason: str | None, |
| candidate_id: str | None, |
| trail: tuple[EvidenceStep, ...], |
| metadata: ModelCallMetadata | None, |
| ) -> VariableTerminalOutcome: |
| return VariableTerminalOutcome( |
| branch_id=branch_id, |
| parent_branch_id=parent_branch_id, |
| topic_uuid=context.topic_uuid, |
| topic_name=context.topic_name, |
| term_uuid=context.term_uuid, |
| term_name=context.term_name, |
| final_uuid=parent.UUID, |
| final_name=parent.name, |
| final_level=parent.level, |
| final_canonical_path=parent.canonical_path, |
| path_components=parent.path_components, |
| evidence=evidence, |
| evidence_trail=trail, |
| support_type=support_type, |
| confidence=confidence, |
| reason=reason, |
| stop_reason=stop_reason, |
| candidate_id=candidate_id, |
| prompt_version=None if metadata is None else metadata.prompt_version, |
| model_provider=None if metadata is None else metadata.provider, |
| model_name=None if metadata is None else metadata.model_name, |
| retry_count=0 if metadata is None else metadata.retry_count, |
| ) |
|
|
|
|
| def _evidence_from_term_branch(term_branch: TermBranchSeed) -> EvidenceStep: |
| return EvidenceStep( |
| branch_id=term_branch.branch_id, |
| UUID=term_branch.term_uuid, |
| name=term_branch.term_name, |
| level=term_branch.term_level, |
| evidence=term_branch.evidence, |
| support_type=term_branch.support_type, |
| confidence=term_branch.confidence, |
| reason=term_branch.reason, |
| candidate_id=term_branch.candidate_id, |
| ) |
|
|
|
|
| def _evidence_from_child(*, child, branch_id: str, decision: CandidateDecision) -> EvidenceStep: |
| return EvidenceStep( |
| branch_id=branch_id, |
| UUID=child.UUID, |
| name=child.name, |
| level=child.level, |
| evidence=decision.evidence, |
| support_type=decision.support_type, |
| confidence=decision.confidence, |
| reason=decision.reason, |
| candidate_id=decision.candidate_id, |
| ) |
|
|
|
|
| def _branch_error( |
| *, |
| branch_id: str, |
| parent_branch_id: str, |
| parent, |
| exc: Exception, |
| ) -> VariableBranchError: |
| return VariableBranchError( |
| branch_id=branch_id, |
| parent_branch_id=parent_branch_id, |
| parent_uuid=parent.UUID, |
| parent_name=parent.name, |
| parent_level=parent.level, |
| code=exc.__class__.__name__, |
| message=str(exc), |
| ) |
|
|
|
|
| def _validate_variable_parent(level: str) -> None: |
| if level not in _VARIABLE_PARENT_LEVELS and level not in _VARIABLE_LEVELS: |
| raise ValueError(f"Unsupported Variable descent parent level: {level!r}.") |
|
|
|
|
| def _model_metadata(response) -> ModelCallMetadata: |
| token_usage = response.token_usage |
| return ModelCallMetadata( |
| provider=response.provider, |
| model_name=response.model_name, |
| prompt_version=response.prompt_version, |
| retry_count=response.retry_count, |
| duration_seconds=response.duration_seconds, |
| input_tokens=None if token_usage is None else token_usage.input_tokens, |
| output_tokens=None if token_usage is None else token_usage.output_tokens, |
| total_tokens=None if token_usage is None else token_usage.total_tokens, |
| estimated_cost=response.estimated_cost, |
| ) |
|
|