flow2 / tests /unit /test_workflows.py
AndrianBalanescu
fix(workflows): survive Qwen3.8 thinking-token exhaustion in chat workflows
0ffaa55
Raw History Blame Contribute Delete
12.8 kB
"""Offline unit tests for moldovan-qwen/workflows OmniRoute workflows.
Uses a scripted fake OmniRouteClient; never touches network or real data
dirs (writes only to pytest tmp_path).
"""
from __future__ import annotations
import json
import math
import os
import sys
from typing import Any, Dict, List
import pytest
sys.path.insert(
0, os.path.join(os.path.dirname(__file__), "..", "..", "moldovan-qwen")
)
from workflows.providers import ( # noqa: E402
EMBEDDING_DIM,
OmniRouteClient,
OmniRouteError,
strip_reasoning,
)
from workflows.summarizer import ( # noqa: E402
_chunk_text,
summarize_document,
)
from workflows.transcriber import transcribe_meeting # noqa: E402
from workflows.semantic_search import ( # noqa: E402
SemanticSearchIndex,
_cosine,
)
from workflows.extractor import ( # noqa: E402
extract_structured_data,
_find_json_block,
)
class FakeClient(OmniRouteClient):
"""Scripted client: records calls, returns queued deterministic outputs."""
def __init__(self, chat_outputs: List[str] = None, embed_fn=None):
super().__init__(base_url="http://fake", api_key="test")
self.chat_outputs = list(chat_outputs or [])
self.chat_calls: List[List[Dict[str, str]]] = []
self.embed_fn = embed_fn
def chat(self, messages, model="1zero", max_tokens=512, temperature=0.7):
self.chat_calls.append(messages)
if not self.chat_outputs:
raise OmniRouteError("no scripted chat outputs left")
return self.chat_outputs.pop(0)
def embed(self, texts, model="infinity/BAAI/bge-m3"):
if self.embed_fn is None:
raise OmniRouteError("no embed_fn configured")
return [self.embed_fn(t) for t in texts]
def transcribe(self, audio_path, model="whisperfw/Systran/faster-whisper-small"):
with open(audio_path, "rb") as f:
header = f.read(12)
if header[:4] != b"RIFF":
raise OmniRouteError("not a wav")
return f"transcribed:{os.path.basename(audio_path)}"
def fake_embed(text: str) -> List[float]:
"""Deterministic 1024-dim vector keyed by first char of text."""
vec = [0.0] * EMBEDDING_DIM
key = sum(ord(c) for c in text.lower()) % EMBEDDING_DIM
vec[key] = 1.0
return vec
# ----------------------------------------------------------------------
# providers
# ----------------------------------------------------------------------
def test_embedding_dim_constant():
assert EMBEDDING_DIM == 1024
def test_strip_reasoning_think_tags():
raw = "<think>internal notes</think>Final answer here."
assert strip_reasoning(raw) == "Final answer here."
def test_strip_reasoning_leaves_plain_text():
assert strip_reasoning("Simple raspuns.") == "Simple raspuns."
def test_strip_reasoning_empty():
assert strip_reasoning("") == ""
assert strip_reasoning(None) is None
def test_strip_reasoning_thinking_process_with_answer():
raw = (
"Thinking Process:\n\n"
"1. **Analyze the Request:**\n"
" * Task: Summarize.\n\n"
"Final summary text: Raspunsul final aici."
)
assert strip_reasoning(raw) == "Final summary text: Raspunsul final aici."
def test_strip_reasoning_truncated_thinking_only():
raw = (
"Thinking Process:\n\n"
"1. **Analyze the Request:**\n"
' * Sentence 2: "Orasul peste 500 mii locuitori" (missing verbs)...'
)
assert strip_reasoning(raw) == ""
def test_chat_error_includes_status(monkeypatch):
client = OmniRouteClient(base_url="http://127.0.0.1:1", api_key="k", timeout=1.0)
import urllib.request
import urllib.error
def raise_http_error(req, timeout):
raise urllib.error.URLError("unreachable")
monkeypatch.setattr(urllib.request, "urlopen", raise_http_error)
with pytest.raises(OmniRouteError):
client.chat([{"role": "user", "content": "hi"}])
# ----------------------------------------------------------------------
# summarizer
# ----------------------------------------------------------------------
def test_chunk_text_short_single_chunk():
assert _chunk_text("Salut, aceasta este o fraza scurta.") == [
"Salut, aceasta este o fraza scurta."
]
def test_chunk_text_empty_returns_empty_list():
assert _chunk_text("") == []
assert _chunk_text(" \n ") == []
def test_chunk_text_long_splits_by_paragraph():
text = "\n\n".join(f"Paragraf {i} " + "x" * 200 for i in range(10))
chunks = _chunk_text(text, max_chars=500)
assert len(chunks) >= 2
assert all(len(c) <= 500 for c in chunks)
joined = " ".join(chunks)
for i in range(10):
assert f"Paragraf {i}" in joined
def test_summarize_single_chunk():
client = FakeClient(chat_outputs=["Rezumat scurt."])
result = summarize_document("Text de test", client)
assert result == {"summary": "Rezumat scurt.", "chunks": 1}
assert len(client.chat_calls) == 1
def test_summarize_multi_chunk_reduces():
long_text = "\n\n".join(
f"Capitol {i} " + "continut lung " * 400 for i in range(4)
)
expected_chunks = len(_chunk_text(long_text))
assert expected_chunks >= 2
# One scripted output per partial chunk, then one final reduce output.
outputs = [f"Partea {i}" for i in range(1, expected_chunks + 1)]
outputs.append("REZUMAT FINAL")
client = FakeClient(chat_outputs=outputs)
result = summarize_document(long_text, client)
assert result["chunks"] == expected_chunks
assert len(client.chat_calls) == expected_chunks + 1
assert result["summary"] == "REZUMAT FINAL"
def test_summarize_empty_raises():
client = FakeClient()
with pytest.raises(ValueError):
summarize_document(" ", client)
def test_summarize_retries_on_empty_content():
class EmptyFirstClient(FakeClient):
def __init__(self):
super().__init__()
self.calls = 0
def chat(self, messages, **kwargs):
self.calls += 1
return "" if self.calls == 1 else "Bun raspuns."
client = EmptyFirstClient()
result = summarize_document("Text scurt de test.", client)
assert result["summary"] == "Bun raspuns."
assert client.calls == 2
def test_summarize_system_prompt_suppresses_thinking():
"""Regression: 1zero leaks 'Thinking Process:' narration and can burn the
whole token budget on reasoning, leaving sanitized content empty. The
summarizer must instruct the model to answer directly without thinking."""
client = FakeClient(chat_outputs=["Rezumat curat."])
summarize_document("Text de test", client)
system_msg = client.chat_calls[0][0]
assert system_msg["role"] == "system"
assert "nu afisa procesul de gandire" in system_msg["content"]
# ----------------------------------------------------------------------
# transcriber
# ----------------------------------------------------------------------
def _write_wav(path):
# minimal fake RIFF header
with open(path, "wb") as f:
f.write(b"RIFF\x00\x00\x00\x00WAVE" + b"\x00" * 16)
return path
def test_transcribe_meeting_single_file(tmp_path):
wav = _write_wav(str(tmp_path / "meeting.wav"))
client = FakeClient()
result = transcribe_meeting([wav], client)
assert result["transcript"] == "transcribed:meeting.wav"
assert result["segments"] == [{"file": "meeting.wav", "text": "transcribed:meeting.wav"}]
def test_transcribe_meeting_directory_sorted(tmp_path):
d = tmp_path / "segs"
d.mkdir()
_write_wav(str(d / "b.wav"))
_write_wav(str(d / "a.wav"))
client = FakeClient()
result = transcribe_meeting([str(d)], client)
assert result["segments"][0]["file"] == "a.wav"
assert result["segments"][1]["file"] == "b.wav"
def test_transcribe_meeting_no_inputs_raises(tmp_path):
client = FakeClient()
with pytest.raises(ValueError):
transcribe_meeting([str(tmp_path / "missing.wav")], client)
with pytest.raises(ValueError):
transcribe_meeting([str(tmp_path)], client) # empty dir
def test_transcribe_meeting_unsupported_ext_raises(tmp_path):
f = tmp_path / "notes.txt"
f.write_text("hello")
client = FakeClient()
with pytest.raises(ValueError):
transcribe_meeting([str(f)], client)
# ----------------------------------------------------------------------
# semantic search
# ----------------------------------------------------------------------
def test_cosine_identical_is_one():
v = [1.0, 2.0, 3.0]
assert math.isclose(_cosine(v, v), 1.0)
def test_cosine_orthogonal_is_zero():
assert _cosine([1.0, 0.0], [0.0, 1.0]) == 0.0
def test_cosine_zero_vector_is_zero():
assert _cosine([0.0, 0.0], [1.0, 1.0]) == 0.0
def test_semantic_search_ranks_relevant_first():
client = FakeClient(embed_fn=fake_embed)
idx = SemanticSearchIndex(client)
idx.add_documents(["mar roshu", "banana galbena", "para verde"])
results = idx.search("banana galbena", top_k=2)
assert len(results) == 2
assert results[0]["document"] == "banana galbena"
assert results[0]["score"] == pytest.approx(1.0)
assert results[0]["index"] == 1
def test_semantic_search_empty_index():
client = FakeClient(embed_fn=fake_embed)
idx = SemanticSearchIndex(client)
assert idx.search("orice") == []
def test_semantic_search_min_score_filter():
client = FakeClient(embed_fn=fake_embed)
idx = SemanticSearchIndex(client)
idx.add_documents(["unu", "doi"])
results = idx.search("unu", top_k=5, min_score=0.99)
assert len(results) == 1
assert results[0]["document"] == "unu"
def test_semantic_index_dim_validation():
client = FakeClient(embed_fn=lambda t: [0.0] * 10) # wrong dim
idx = SemanticSearchIndex(client)
with pytest.raises(ValueError):
idx.add_documents(["text"])
def test_semantic_index_save_load_roundtrip(tmp_path):
client = FakeClient(embed_fn=fake_embed)
idx = SemanticSearchIndex(client)
idx.add_documents(["mar", "para", "pruna"])
path = str(tmp_path / "index.json")
idx.save(path)
loaded = SemanticSearchIndex.load(path, client)
assert loaded.documents == idx.documents
assert loaded.vectors == idx.vectors
assert loaded.search("para")[0]["document"] == "para"
def test_semantic_index_load_dim_mismatch(tmp_path):
path = str(tmp_path / "bad.json")
with open(path, "w") as f:
json.dump({"model_dim": 768, "documents": [], "vectors": []}, f)
with pytest.raises(ValueError):
SemanticSearchIndex.load(path)
def test_semantic_index_load_len_mismatch(tmp_path):
client = FakeClient(embed_fn=fake_embed)
path = str(tmp_path / "bad2.json")
with open(path, "w") as f:
json.dump(
{"model_dim": 1024, "documents": ["a"], "vectors": []}, f
)
with pytest.raises(ValueError):
SemanticSearchIndex.load(path, client)
# ----------------------------------------------------------------------
# extractor
# ----------------------------------------------------------------------
def test_find_json_block_fenced():
assert _find_json_block("text ```json\n{\"a\": 1}\n``` end") == '{"a": 1}'
def test_find_json_block_plain_braces():
assert _find_json_block('Iată: {"a": 2} gata') == '{"a": 2}'
def test_find_json_block_none():
assert _find_json_block("niciun json aici") is None
def test_extract_structured_data_success():
client = FakeClient(chat_outputs=['{"nume": "Ion", "varsta": 40}'])
schema = {"nume": "string", "varsta": "number"}
result = extract_structured_data("Ion are 40 de ani.", schema, client)
assert result == {"nume": "Ion", "varsta": 40}
def test_extract_structured_data_retry_on_garbage():
client = FakeClient(
chat_outputs=[
"Nu pot raspunde", # first attempt garbage
'{"oras": "Chisinau"}', # retry succeeds
]
)
result = extract_structured_data("Locuiesc in Chisinau.", {"oras": "string"}, client)
assert result == {"oras": "Chisinau"}
assert len(client.chat_calls) == 2 # original + retry
def test_extract_structured_data_fails_after_retry():
client = FakeClient(chat_outputs=["garbage 1", "garbage 2"])
with pytest.raises(OmniRouteError):
extract_structured_data("text", {"a": "string"}, client)
def test_extract_structured_data_empty_text_raises():
client = FakeClient()
with pytest.raises(ValueError):
extract_structured_data("", {"a": "string"}, client)
def test_extract_structured_data_with_fenced_json():
client = FakeClient(chat_outputs=["```json\n{\"tara\": \"MD\"}\n```"])
result = extract_structured_data("Republica Moldova", {"tara": "string"}, client)
assert result == {"tara": "MD"}