Spaces:
Running on Zero
Running on Zero
AndrianBalanescu
fix(workflows): survive Qwen3.8 thinking-token exhaustion in chat workflows
0ffaa55 Download tests/unit/test_workflows.py from abalanescu/flow2: direct link, hf CLI and curl.
- Browser
- Download file 12.8 kB
-
https://huggingface.co/spaces/abalanescu/flow2/resolve/main/tests/unit/test_workflows.py
- Command line
-
hf download hf://spaces/abalanescu/flow2/tests/unit/test_workflows.py
-
curl -L -o test_workflows.py https://huggingface.co/spaces/abalanescu/flow2/resolve/main/tests/unit/test_workflows.py
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"} |