File size: 8,000 Bytes
9f2f534 08942f5 9f2f534 08942f5 9f2f534 | 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 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 | from __future__ import annotations
import json
import os
import sys
from pathlib import Path
import pytest
from jsonschema import Draft202012Validator
from gcmd_classifier.articles import load_articles
from gcmd_classifier.config import ModelSettings
from gcmd_classifier.llm import FakeModelClient
from gcmd_classifier.llm.openai_provider import OpenAIModelClient
from gcmd_classifier.models import (
ArticleClassificationOutcome,
ArticleLoadResult,
ArticleProcessingStatus,
)
from gcmd_classifier.persistence import ArticleResultCache, JsonResultStore
from gcmd_classifier.pipeline import run_batch
from gcmd_classifier.vocabulary import VocabularyIndex, load_vocabulary
FULL_HIERARCHY_PATH = Path("data/gcmd_hierarchy.json")
FULL_ARTICLES_PATH = Path("data/articles.json")
PROTOTYPE_PATH = Path("prototype/app_hf_poc.py")
CLASSIFICATION_SCHEMA_PATH = Path("schemas/classification_result.schema.json")
RUN_SUMMARY_SCHEMA_PATH = Path("schemas/run_summary.schema.json")
def _decision(candidate_id: str) -> dict:
return {
"candidate_id": candidate_id,
"confidence": 0.9,
"evidence": f"Fake smoke evidence for {candidate_id}.",
"support_type": "explicit",
"reason": f"Fake smoke selected {candidate_id}.",
}
def _topic_response(candidate_id: str) -> dict:
return {"selected": [_decision(candidate_id)]}
def _term_response(candidate_id: str) -> dict:
return {"selected": [_decision(candidate_id)], "stop_at_parent": False}
def _variable_response(candidate_id: str = "variable_0001") -> dict:
return {"selected": [_decision(candidate_id)], "stop_at_parent": False}
def _no_topic_response() -> dict:
return {
"selected": [],
"no_selection_reason": "Fake smoke no-classification article.",
}
def _fake_smoke_actions(index: VocabularyIndex) -> list[dict]:
topic_position, topic, term_position, term = _first_term_with_variables(index)
actions = [
_topic_response(f"topic_{topic_position:04d}"),
_term_response(f"term_{term_position:04d}"),
]
parent = term
while parent.level != "Variable_Level_3":
children = index.variables_for_parent(parent.UUID)
if not children:
break
actions.append(_variable_response("variable_0001"))
parent = children[0]
actions.append(_no_topic_response())
return actions
def _first_term_with_variables(index: VocabularyIndex):
for topic_position, topic in enumerate(index.topics(), start=1):
for term_position, term in enumerate(index.terms_for_topic(topic.UUID), start=1):
if index.variables_for_parent(term.UUID):
return topic_position, topic, term_position, term
raise AssertionError("Current vocabulary did not contain a Term with Variable children.")
def _smoke_article_load_result(full_load: ArticleLoadResult) -> ArticleLoadResult:
return full_load.model_copy(update={"articles": full_load.articles[:2]})
def test_fake_model_mvp_smoke_run_end_to_end(tmp_path: Path) -> None:
hierarchy_before = FULL_HIERARCHY_PATH.read_bytes()
articles_before = FULL_ARTICLES_PATH.read_bytes()
prototype_before = PROTOTYPE_PATH.read_bytes()
classification_schema = json.loads(CLASSIFICATION_SCHEMA_PATH.read_text())
summary_schema = json.loads(RUN_SUMMARY_SCHEMA_PATH.read_text())
vocabulary = load_vocabulary(FULL_HIERARCHY_PATH)
raw_records = json.loads(FULL_ARTICLES_PATH.read_text())
full_load = load_articles(FULL_ARTICLES_PATH)
smoke_load = _smoke_article_load_result(full_load)
cache = ArticleResultCache(tmp_path / "cache")
store = JsonResultStore(tmp_path / "results")
batch = run_batch(
article_load_result=smoke_load,
vocabulary=vocabulary,
model_client=FakeModelClient(_fake_smoke_actions(vocabulary)),
settings=ModelSettings(),
store=store,
cache=cache,
run_id="fake-smoke",
relevant_config={"smoke_subset_size": len(smoke_load.articles)},
)
assert len(vocabulary) == 3535
assert isinstance(raw_records, list)
assert full_load.source_count == len(raw_records)
assert len(full_load.articles) == len(raw_records) - len(full_load.errors)
assert len({article.DOI for article in full_load.articles}) == len(full_load.articles)
if full_load.errors:
assert all(error.code and error.message for error in full_load.errors)
assert all(
error.index is None or 0 <= error.index < len(raw_records) for error in full_load.errors
)
assert batch.summary.articles_received == len(raw_records)
assert batch.summary.invalid_source_records == len(full_load.errors)
assert batch.summary.processed_articles == 2
assert batch.summary.cache_misses == 2
assert batch.summary.cache_hits == 0
assert batch.summary.total_model_calls and batch.summary.total_model_calls > 0
assert batch.summary.duration_seconds is not None
assert batch.summary.started_at is not None
assert batch.summary.completed_at is not None
classified = [
result
for result in batch.results
if result.classification_outcome is ArticleClassificationOutcome.CLASSIFIED
]
not_classified = [
result
for result in batch.results
if result.classification_outcome is ArticleClassificationOutcome.NOT_CLASSIFIED
]
assert classified
assert not_classified
assert not_classified[0].processing_status is ArticleProcessingStatus.COMPLETED
assert not_classified[0].classifications == ()
assert not_classified[0].no_classification_reason
accepted = tuple(record for result in batch.results for record in result.classifications)
assert accepted
assert batch.summary.accepted_classifications == len(accepted)
assert all(record.deterministic_validation.valid for record in accepted)
assert not any(
record.final_status == "accepted" and not record.deterministic_validation.valid
for record in accepted
)
article_validator = Draft202012Validator(classification_schema)
for result in batch.results:
article_validator.validate(result.model_dump(mode="json"))
Draft202012Validator(summary_schema).validate(batch.summary.model_dump(mode="json"))
assert store.consolidated_path.exists()
persisted = json.loads(store.consolidated_path.read_text())
assert persisted["summary"]["run_id"] == "fake-smoke"
assert len(persisted["articles"]) == 2
assert "prototype.app_hf_poc" not in sys.modules
assert FULL_HIERARCHY_PATH.read_bytes() == hierarchy_before
assert FULL_ARTICLES_PATH.read_bytes() == articles_before
assert PROTOTYPE_PATH.read_bytes() == prototype_before
@pytest.mark.integration
def test_optional_live_model_smoke_is_skipped_unless_explicitly_enabled(tmp_path: Path) -> None:
if os.environ.get("GCMD_RUN_LIVE_INTEGRATION") != "1":
pytest.skip("Set GCMD_RUN_LIVE_INTEGRATION=1 to enable live model smoke checks.")
settings = ModelSettings.from_environment()
if settings.provider != "openai":
pytest.skip("Set MODEL_PROVIDER=openai for live model smoke checks.")
if not os.environ.get(settings.api_key_env_var):
pytest.skip(f"Set {settings.api_key_env_var} for live model smoke checks.")
vocabulary = load_vocabulary(FULL_HIERARCHY_PATH)
full_load = load_articles(FULL_ARTICLES_PATH)
smoke_load = full_load.model_copy(update={"articles": full_load.articles[:1]})
batch = run_batch(
article_load_result=smoke_load,
vocabulary=vocabulary,
model_client=OpenAIModelClient(settings),
settings=settings,
store=JsonResultStore(tmp_path / "live-results"),
run_id="live-smoke",
)
accepted = tuple(record for result in batch.results for record in result.classifications)
assert all(record.deterministic_validation.valid for record in accepted)
assert batch.summary.processed_articles == 1
|