| from __future__ import annotations |
|
|
| import io |
| import json |
| import os |
| import subprocess |
| import sys |
|
|
| from deberta_ime import RerankRequest |
| from deberta_ime.cli import run |
| from deberta_ime.mozc import MOZC_REVISION |
|
|
|
|
| class ContextScorer: |
| def score_candidates(self, request: RerankRequest) -> list[float]: |
| return [-10.0, 0.0] |
|
|
|
|
| class RecordingScorer: |
| def __init__(self) -> None: |
| self.calls = 0 |
|
|
| def score_candidates(self, request: RerankRequest) -> list[float]: |
| self.calls += 1 |
| return [0.0] * len(request.candidates) |
|
|
|
|
| class LoadableScorer: |
| def __init__(self) -> None: |
| self.loaded = False |
|
|
| def load(self) -> None: |
| self.loaded = True |
|
|
| def score_candidates(self, request: RerankRequest) -> list[float]: |
| if not self.loaded: |
| raise RuntimeError("not warmed") |
| return [0.0] * len(request.candidates) |
|
|
|
|
| class BrokenLoadScorer: |
| def load(self) -> None: |
| raise RuntimeError("model cache unavailable") |
|
|
| def score_candidates(self, request: RerankRequest) -> list[float]: |
| raise AssertionError("scoring should not run") |
|
|
|
|
| def test_server_echoes_request_id_and_returns_only_original_candidate_ranks() -> None: |
| request = { |
| "schema_version": 1, |
| "request_id": "host-橋-1", |
| "op": "rerank", |
| "candidate_source": { |
| "id": "mozc_oss_dictionary_cost_v1", |
| "revision": MOZC_REVISION, |
| }, |
| "mode": "bidirectional", |
| "reading": "はし", |
| "left_context": ["川", "に", "架かる"], |
| "right_context": ["を", "渡る"], |
| "candidates": [ |
| {"surface": "箸", "prior_score": 0.0}, |
| {"surface": "橋", "prior_score": -0.3}, |
| ], |
| } |
| stdout = io.StringIO() |
|
|
| exit_code = run( |
| ["serve"], |
| stdin=io.StringIO(json.dumps(request, ensure_ascii=False) + "\n"), |
| stdout=stdout, |
| scorer=ContextScorer(), |
| ) |
|
|
| response = json.loads(stdout.getvalue()) |
| assert exit_code == 0 |
| assert response["ok"] is True |
| assert response["schema_version"] == 1 |
| assert response["request_id"] == "host-橋-1" |
| assert response["ranked_original_ranks"] == [1, 0] |
| assert response["decision"] == { |
| "changed": True, |
| "reason": "accepted", |
| "margin": 9.1, |
| } |
| assert "candidates" not in response |
|
|
|
|
| def test_server_preserves_order_without_scoring_an_uncalibrated_source() -> None: |
| scorer = RecordingScorer() |
| request = { |
| "schema_version": 1, |
| "request_id": "unknown-source-1", |
| "op": "rerank", |
| "candidate_source": {"id": "microsoft_ime_display_rank", "revision": "unknown"}, |
| "mode": "incremental", |
| "reading": "はし", |
| "left_context": ["川"], |
| "right_context": [], |
| "candidates": [ |
| {"surface": "箸", "prior_score": 0.0}, |
| {"surface": "橋", "prior_score": -1.0}, |
| ], |
| } |
| stdout = io.StringIO() |
|
|
| exit_code = run( |
| ["serve"], |
| stdin=io.StringIO(json.dumps(request, ensure_ascii=False) + "\n"), |
| stdout=stdout, |
| scorer=scorer, |
| ) |
|
|
| response = json.loads(stdout.getvalue()) |
| assert exit_code == 0 |
| assert response["ok"] is True |
| assert response["request_id"] == "unknown-source-1" |
| assert response["ranked_original_ranks"] == [0, 1] |
| assert response["decision"] == { |
| "changed": False, |
| "reason": "uncalibrated_source", |
| "margin": None, |
| } |
| assert response["profile"] is None |
| assert scorer.calls == 0 |
|
|
|
|
| def test_incremental_server_request_preserves_order_when_future_context_is_present() -> None: |
| scorer = RecordingScorer() |
| request = { |
| "schema_version": 1, |
| "request_id": "context-leak-1", |
| "op": "rerank", |
| "candidate_source": { |
| "id": "mozc_oss_dictionary_cost_v1", |
| "revision": MOZC_REVISION, |
| }, |
| "mode": "incremental", |
| "reading": "はし", |
| "left_context": ["川"], |
| "right_context": ["を", "渡る"], |
| "candidates": [ |
| {"surface": "箸", "prior_score": 0.0}, |
| {"surface": "橋", "prior_score": -0.3}, |
| ], |
| } |
| stdout = io.StringIO() |
|
|
| run( |
| ["serve"], |
| stdin=io.StringIO(json.dumps(request, ensure_ascii=False) + "\n"), |
| stdout=stdout, |
| scorer=scorer, |
| ) |
|
|
| response = json.loads(stdout.getvalue()) |
| assert response["ranked_original_ranks"] == [0, 1] |
| assert response["decision"] == { |
| "changed": False, |
| "reason": "context_mode_mismatch", |
| "margin": None, |
| } |
| assert scorer.calls == 0 |
|
|
|
|
| def test_health_is_cold_and_warmup_reports_model_ready_in_the_same_process() -> None: |
| requests = [ |
| {"schema_version": 1, "request_id": "health-1", "op": "health"}, |
| {"schema_version": 1, "request_id": "warmup-1", "op": "warmup"}, |
| ] |
| stdin = io.StringIO( |
| "".join(json.dumps(item, ensure_ascii=False) + "\n" for item in requests) |
| ) |
| stdout = io.StringIO() |
|
|
| exit_code = run(["serve"], stdin=stdin, stdout=stdout, scorer=LoadableScorer()) |
|
|
| responses = [json.loads(line) for line in stdout.getvalue().splitlines()] |
| assert exit_code == 0 |
| assert responses[0] == { |
| "ok": True, |
| "schema_version": 1, |
| "request_id": "health-1", |
| "operation": "health", |
| "model_loaded": False, |
| } |
| assert responses[1] == { |
| "ok": True, |
| "schema_version": 1, |
| "request_id": "warmup-1", |
| "operation": "warmup", |
| "model_loaded": True, |
| } |
|
|
|
|
| def test_unsupported_context_mode_preserves_order_without_scoring() -> None: |
| scorer = RecordingScorer() |
| request = { |
| "schema_version": 1, |
| "request_id": "mode-1", |
| "op": "rerank", |
| "candidate_source": { |
| "id": "mozc_oss_dictionary_cost_v1", |
| "revision": MOZC_REVISION, |
| }, |
| "mode": "future_context_guessing", |
| "reading": "はし", |
| "candidates": [ |
| {"surface": "箸", "prior_score": 0.0}, |
| {"surface": "橋", "prior_score": -0.3}, |
| ], |
| } |
| stdout = io.StringIO() |
|
|
| run( |
| ["serve"], |
| stdin=io.StringIO(json.dumps(request, ensure_ascii=False) + "\n"), |
| stdout=stdout, |
| scorer=scorer, |
| ) |
|
|
| response = json.loads(stdout.getvalue()) |
| assert response["ok"] is True |
| assert response["ranked_original_ranks"] == [0, 1] |
| assert response["decision"]["reason"] == "unsupported_mode" |
| assert scorer.calls == 0 |
|
|
|
|
| def test_real_server_protocol_stays_utf8_when_console_encoding_is_cp932() -> None: |
| request = { |
| "schema_version": 1, |
| "request_id": "utf8-橋-🧪", |
| "op": "health", |
| } |
| environment = os.environ.copy() |
| environment["PYTHONIOENCODING"] = "cp932" |
|
|
| completed = subprocess.run( |
| [sys.executable, "-m", "deberta_ime", "serve", "--offline"], |
| input=(json.dumps(request, ensure_ascii=False) + "\n").encode("utf-8"), |
| capture_output=True, |
| env=environment, |
| check=False, |
| timeout=30, |
| ) |
|
|
| assert completed.returncode == 0, completed.stderr.decode("utf-8", errors="replace") |
| response = json.loads(completed.stdout.decode("utf-8")) |
| assert response["request_id"] == "utf8-橋-🧪" |
| assert response["operation"] == "health" |
|
|
|
|
| def test_boolean_prior_is_not_accepted_as_a_numeric_candidate_score() -> None: |
| scorer = RecordingScorer() |
| request = { |
| "schema_version": 1, |
| "request_id": "bool-prior-1", |
| "op": "rerank", |
| "candidate_source": { |
| "id": "mozc_oss_dictionary_cost_v1", |
| "revision": MOZC_REVISION, |
| }, |
| "mode": "incremental", |
| "reading": "はし", |
| "candidates": [ |
| {"surface": "箸", "prior_score": True}, |
| {"surface": "橋", "prior_score": -0.3}, |
| ], |
| } |
| stdout = io.StringIO() |
|
|
| run( |
| ["serve"], |
| stdin=io.StringIO(json.dumps(request, ensure_ascii=False) + "\n"), |
| stdout=stdout, |
| scorer=scorer, |
| ) |
|
|
| response = json.loads(stdout.getvalue()) |
| assert response["ok"] is False |
| assert response["request_id"] == "bool-prior-1" |
| assert scorer.calls == 0 |
|
|
|
|
| def test_oversized_request_line_is_rejected_before_model_scoring() -> None: |
| scorer = RecordingScorer() |
| request = { |
| "schema_version": 1, |
| "request_id": "large-line-1", |
| "op": "rerank", |
| "candidate_source": { |
| "id": "mozc_oss_dictionary_cost_v1", |
| "revision": MOZC_REVISION, |
| }, |
| "mode": "incremental", |
| "reading": "はし", |
| "candidates": [ |
| {"surface": "箸", "prior_score": 0.0}, |
| {"surface": "橋", "prior_score": -0.3}, |
| ], |
| "ignored_padding": "x" * 70_000, |
| } |
| stdout = io.StringIO() |
|
|
| run( |
| ["serve"], |
| stdin=io.StringIO(json.dumps(request, ensure_ascii=False) + "\n"), |
| stdout=stdout, |
| scorer=scorer, |
| ) |
|
|
| response = json.loads(stdout.getvalue()) |
| assert response["ok"] is False |
| assert response["request_id"] == "large-line-1" |
| assert response["error"] == "request line exceeds 65536 UTF-8 bytes" |
| assert scorer.calls == 0 |
|
|
|
|
| def test_model_warmup_failure_does_not_terminate_the_protocol_loop() -> None: |
| requests = [ |
| {"schema_version": 1, "request_id": "warmup-fail", "op": "warmup"}, |
| {"schema_version": 1, "request_id": "health-after", "op": "health"}, |
| ] |
| stdout = io.StringIO() |
|
|
| run( |
| ["serve"], |
| stdin=io.StringIO("".join(json.dumps(item) + "\n" for item in requests)), |
| stdout=stdout, |
| scorer=BrokenLoadScorer(), |
| ) |
|
|
| responses = [json.loads(line) for line in stdout.getvalue().splitlines()] |
| assert responses[0]["ok"] is False |
| assert responses[0]["request_id"] == "warmup-fail" |
| assert responses[1] == { |
| "ok": True, |
| "schema_version": 1, |
| "request_id": "health-after", |
| "operation": "health", |
| "model_loaded": False, |
| } |
|
|
|
|
| def test_overlong_reading_is_preserved_without_scoring() -> None: |
| scorer = RecordingScorer() |
| request = { |
| "schema_version": 1, |
| "request_id": "reading-limit-1", |
| "op": "rerank", |
| "candidate_source": { |
| "id": "mozc_oss_dictionary_cost_v1", |
| "revision": MOZC_REVISION, |
| }, |
| "mode": "incremental", |
| "reading": "あ" * 129, |
| "candidates": [ |
| {"surface": "亜", "prior_score": 0.0}, |
| {"surface": "阿", "prior_score": -0.1}, |
| ], |
| } |
| stdout = io.StringIO() |
|
|
| run( |
| ["serve"], |
| stdin=io.StringIO(json.dumps(request, ensure_ascii=False) + "\n"), |
| stdout=stdout, |
| scorer=scorer, |
| ) |
|
|
| response = json.loads(stdout.getvalue()) |
| assert response["ranked_original_ranks"] == [0, 1] |
| assert response["decision"]["reason"] == "unsafe_input" |
| assert scorer.calls == 0 |
|
|