Team Ai
Apppublic

KillerKing93/Transformers-InferenceServer-OpenAPI

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes
test_api.py712 linesDownload Raw Back to tests
1import json2import time3from contextlib import contextmanager4 5import pytest6from fastapi.testclient import TestClient7 8import main9 10 11class FakeEngine:12    def __init__(self, model_id="fake-model"):13        self.model_id = model_id14        self.last_context_info = {15            "compressed": False,16            "prompt_tokens": 5,17            "max_context": 8192,18            "budget": 7900,19            "strategy": "truncate",20            "dropped_messages": 0,21        }22 23    def infer(self, messages, max_tokens, temperature):24        # Simulate parse error pathway when special trigger is present25        if messages and isinstance(messages[0].get("content"), str) and "PARSE_ERR" in messages[0]["content"]:26            raise ValueError("Simulated parse error")27        # Return echo content for deterministic test28        parts = []29        for m in messages:30            c = m.get("content", "")31            if isinstance(c, list):32                for p in c:33                    if isinstance(p, dict) and p.get("type") == "text":34                        parts.append(p.get("text", ""))35            elif isinstance(c, str):36                parts.append(c)37        txt = " ".join(parts) or "OK"38        # Simulate context accounting changing with request39        self.last_context_info = {40            "compressed": False,41            "prompt_tokens": max(1, len(txt.split())),42            "max_context": 8192,43            "budget": 7900,44            "strategy": "truncate",45            "dropped_messages": 0,46        }47        return f"OK: {txt}"48 49    def infer_stream(self, messages, max_tokens, temperature, cancel_event=None):50        # simple two-piece stream; respects cancel_event if set during streaming51        outputs = ["hello", " world"]52        for piece in outputs:53            if cancel_event is not None and cancel_event.is_set():54                break55            yield piece56            # tiny delay to allow cancel test to interleave57            time.sleep(0.01)58 59    def get_context_report(self):60        return {61            "compressionEnabled": True,62            "strategy": "truncate",63            "safetyMargin": 256,64            "modelMaxContext": 8192,65            "tokenizerModelMaxLength": 8192,66            "last": self.last_context_info,67        }68 69 70@contextmanager71def patched_engine():72    # Patch global engine so server does not load real model73    prev_engine = main._engine74    prev_err = main._engine_error75    fake = FakeEngine()76    main._engine = fake77    main._engine_error = None78    try:79        yield fake80    finally:81        main._engine = prev_engine82        main._engine_error = prev_err83 84 85def get_client():86    return TestClient(main.app)87 88 89def test_health_ready_and_context():90    with patched_engine():91        client = get_client()92        r = client.get("/health")93        assert r.status_code == 20094        body = r.json()95        assert body["ok"] is True96        assert body["modelReady"] is True97        assert body["modelId"] == "fake-model"98        # context block exists with required fields99        ctx = body["context"]100        assert ctx["compressionEnabled"] is True101        assert "last" in ctx102        assert isinstance(ctx["last"].get("prompt_tokens"), int)103 104 105def test_health_with_engine_error():106    # simulate model load error path107    prev_engine = main._engine108    prev_err = main._engine_error109    try:110        main._engine = None111        main._engine_error = "boom"112        client = get_client()113        r = client.get("/health")114        assert r.status_code == 200115        body = r.json()116        assert body["modelReady"] is False117        assert body["error"] == "boom"118    finally:119        main._engine = prev_engine120        main._engine_error = prev_err121 122 123def test_chat_non_stream_validation():124    with patched_engine():125        client = get_client()126        # missing messages should 400127        r = client.post("/v1/chat/completions", json={"messages": []})128        assert r.status_code == 400129 130 131def test_chat_non_stream_success_and_usage_context():132    with patched_engine():133        client = get_client()134        payload = {135            "messages": [{"role": "user", "content": "Hello Qwen"}],136            "max_tokens": 8,137            "temperature": 0.0,138        }139        r = client.post("/v1/chat/completions", json=payload)140        assert r.status_code == 200141        body = r.json()142        assert body["object"] == "chat.completion"143        assert body["choices"][0]["message"]["content"].startswith("OK:")144        # usage prompt_tokens filled from engine.last_context_info145        assert body["usage"]["prompt_tokens"] >= 1146        # response includes context echo147        assert "context" in body148        assert "prompt_tokens" in body["context"]149 150 151def test_chat_non_stream_parse_error_to_400():152    with patched_engine():153        client = get_client()154        payload = {155            "messages": [{"role": "user", "content": "PARSE_ERR trigger"}],156            "max_tokens": 4,157        }158        r = client.post("/v1/chat/completions", json=payload)159        # ValueError in engine -> 400 per API contract160        assert r.status_code == 400161 162 163def read_sse_lines(resp):164    # Utility to parse event-stream into list of data payloads (including [DONE])165    lines = []166    buf = b""167 168    # Starlette TestClient (httpx) responses expose iter_bytes()/iter_raw(), not requests.iter_content().169    # Fall back to available iterator or to full content if streaming isn't supported.170    iterator = None171    for name in ("iter_bytes", "iter_raw", "iter_content"):172        it = getattr(resp, name, None)173        if callable(it):174            iterator = it175            break176 177    if iterator is None:178        data = getattr(resp, "content", b"")179        if isinstance(data, str):180            data = data.encode("utf-8", "ignore")181        buf = data182    else:183        for chunk in iterator():184            if not chunk:185                continue186            if isinstance(chunk, str):187                chunk = chunk.encode("utf-8", "ignore")188            buf += chunk189            while b"\n\n" in buf:190                frame, buf = buf.split(b"\n\n", 1)191                # keep original frame text for asserts192                lines.append(frame.decode("utf-8", errors="ignore"))193 194    # Drain any leftover195    if buf:196        lines.append(buf.decode("utf-8", errors="ignore"))197    return lines198 199 200def test_chat_stream_sse_flow_and_resume():201    with patched_engine():202        client = get_client()203        payload = {204            "session_id": "s1",205            "stream": True,206            "messages": [{"role": "user", "content": "stream please"}],207            "max_tokens": 8,208            "temperature": 0.2,209        }210        with client.stream("POST", "/v1/chat/completions", json=payload) as resp:211            assert resp.status_code == 200212            lines = read_sse_lines(resp)213        # Must contain role delta, content pieces, finish chunk, and [DONE]214        joined = "\n".join(lines)215        assert "delta" in joined216        assert "[DONE]" in joined217 218        # Resume from event index 0 should receive at least one subsequent event219        headers = {"Last-Event-ID": "s1:0"}220        with client.stream("POST", "/v1/chat/completions", headers=headers, json=payload) as resp2:221            assert resp2.status_code == 200222            lines2 = read_sse_lines(resp2)223        assert any("data:" in l for l in lines2)224        assert "[DONE]" in "\n".join(lines2)225 226        # Invalid Last-Event-ID format should not crash (covered by try/except)227        headers_bad = {"Last-Event-ID": "not-an-index"}228        with client.stream("POST", "/v1/chat/completions", headers=headers_bad, json=payload) as resp3:229            assert resp3.status_code == 200230            _ = read_sse_lines(resp3)  # just ensure no crash231 232 233def test_cancel_endpoint_stops_generation():234    with patched_engine():235        client = get_client()236        payload = {237            "session_id": "to-cancel",238            "stream": True,239            "messages": [{"role": "user", "content": "cancel me"}],240        }241        # Start streaming in background (client.stream keeps the connection open)242        with client.stream("POST", "/v1/chat/completions", json=payload) as resp:243            # Immediately cancel244            rc = client.post("/v1/cancel/to-cancel")245            assert rc.status_code == 200246            # Stream should end with [DONE] without hanging247            lines = read_sse_lines(resp)248            assert "[DONE]" in "\n".join(lines)249 250 251def test_cancel_unknown_session_is_ok():252    with patched_engine():253        client = get_client()254        rc = client.post("/v1/cancel/does-not-exist")255        # Endpoint returns ok regardless (idempotent, operationally safe)256        assert rc.status_code == 200257 258 259def test_edge_large_last_event_id_after_finish_yields_done():260    with patched_engine():261        client = get_client()262        payload = {263            "session_id": "done-session",264            "stream": True,265            "messages": [{"role": "user", "content": "edge"}],266        }267        # Complete a run268        with client.stream("POST", "/v1/chat/completions", json=payload) as resp:269            _ = read_sse_lines(resp)270        # Resume with huge index; should return DONE quickly271        headers = {"Last-Event-ID": "done-session:99999"}272        with client.stream("POST", "/v1/chat/completions", headers=headers, json=payload) as resp2:273            lines2 = read_sse_lines(resp2)274        assert "[DONE]" in "\n".join(lines2)275 276 277def test_stream_resume_basic_functionality():278    """Test that basic streaming resume functionality works correctly"""279    with patched_engine():280        client = get_client()281        session_id = "resume-basic-test"282        payload = {283            "session_id": session_id,284            "stream": True,285            "messages": [{"role": "user", "content": "test resume"}],286            "max_tokens": 100,287        }288 289        # Complete a streaming session first290        with client.stream("POST", "/v1/chat/completions", json=payload) as resp:291            complete_lines = read_sse_lines(resp)292 293        # Verify we got some data lines294        data_lines = [line for line in complete_lines if line.startswith("data: ")]295        assert len(data_lines) > 0, "Should have received some data lines in complete session"296 297        # Test resume from index 0 (should replay everything)298        headers = {"Last-Event-ID": f"{session_id}:0"}299        with client.stream("POST", "/v1/chat/completions", headers=headers, json=payload) as resp2:300            resume_lines = read_sse_lines(resp2)301 302        # Should get data lines again303        resume_data_lines = [line for line in resume_lines if line.startswith("data: ")]304        assert len(resume_data_lines) > 0, "Should have received data lines on resume"305 306        # Should end with [DONE]307        assert any("[DONE]" in line for line in resume_lines), "Resume should end with [DONE]"308 309 310def test_stream_resume_preserves_exact_chunk_order():311    """Test that resume maintains exact chunk order and content"""312    with patched_engine():313        client = get_client()314        session_id = "order-test"315        payload = {316            "session_id": session_id,317            "stream": True,318            "messages": [{"role": "user", "content": "test order"}],319        }320 321        # Get complete session with default chunks322        with client.stream("POST", "/v1/chat/completions", json=payload) as resp:323            complete_lines = read_sse_lines(resp)324 325        complete_chunks = []326        for line in complete_lines:327            if line.startswith("data: "):328                try:329                    data = json.loads(line[6:])330                    if "choices" in data and data["choices"][0].get("delta", {}).get("content"):331                        complete_chunks.append(data["choices"][0]["delta"]["content"])332                except (json.JSONDecodeError, KeyError):333                    continue334 335        # Skip test if no chunks received (default FakeEngine may not produce content chunks)336        if len(complete_chunks) == 0:337            pytest.skip("Default FakeEngine does not produce content chunks for this test")338 339        # Test resume from middle point340        resume_point = min(1, len(complete_chunks) - 1)  # Resume after first chunk, or 0 if only 1 chunk341        headers = {"Last-Event-ID": f"{session_id}:{resume_point}"}342 343        with client.stream("POST", "/v1/chat/completions", headers=headers, json=payload) as resp:344            resume_lines = read_sse_lines(resp)345 346        resume_chunks = []347        for line in resume_lines:348            if line.startswith("data: "):349                try:350                    data = json.loads(line[6:])351                    if "choices" in data and data["choices"][0].get("delta", {}).get("content"):352                        resume_chunks.append(data["choices"][0]["delta"]["content"])353                except (json.JSONDecodeError, KeyError):354                    continue355 356        # Should get remaining chunks after resume point357        expected_resume_chunks = complete_chunks[resume_point:]358 359        assert resume_chunks == expected_resume_chunks, (360            f"Order preservation failed. Resume chunks: {resume_chunks}, "361            f"Expected: {expected_resume_chunks}"362        )363 364 365def test_stream_resume_with_partial_disconnect():366    """Test resume when client disconnects mid-stream and reconnects"""367    with patched_engine():368        client = get_client()369        session_id = "disconnect-test"370        payload = {371            "session_id": session_id,372            "stream": True,373            "messages": [{"role": "user", "content": "test disconnect"}],374        }375 376        # Complete a session first to populate buffers377        with client.stream("POST", "/v1/chat/completions", json=payload) as resp:378            complete_lines = read_sse_lines(resp)379 380        # Extract chunks and event IDs381        received_chunks = []382        event_ids = []383        for line in complete_lines:384            if line.startswith("data: "):385                try:386                    data = json.loads(line[6:])387                    if "choices" in data and data["choices"][0].get("delta", {}).get("content"):388                        chunk = data["choices"][0]["delta"]["content"]389                        received_chunks.append(chunk)390                        if "id" in data:391                            event_ids.append(data["id"])392                except (json.JSONDecodeError, KeyError):393                    continue394 395        if len(received_chunks) < 2:396            pytest.skip("Not enough chunks received for disconnect test")397 398        # Simulate disconnect after first chunk399        resume_index = 0  # Resume from beginning400        headers = {"Last-Event-ID": f"{session_id}:{resume_index}"}401 402        with client.stream("POST", "/v1/chat/completions", headers=headers, json=payload) as resp:403            resume_lines = read_sse_lines(resp)404 405        resume_chunks = []406        for line in resume_lines:407            if line.startswith("data: "):408                try:409                    data = json.loads(line[6:])410                    if "choices" in data and data["choices"][0].get("delta", {}).get("content"):411                        resume_chunks.append(data["choices"][0]["delta"]["content"])412                except (json.JSONDecodeError, KeyError):413                    continue414 415        # Should get all chunks again (resume from 0 replays everything)416        assert resume_chunks == received_chunks, (417            f"Resume from beginning failed. Got: {resume_chunks}, Expected: {received_chunks}"418        )419 420 421def test_stream_resume_buffer_overflow_handling():422    """Test resume when buffer overflows and older chunks are lost"""423    with patched_engine():424        client = get_client()425        session_id = "overflow-test"426        payload = {427            "session_id": session_id,428            "stream": True,429            "messages": [{"role": "user", "content": "test overflow"}],430        }431 432        # Complete a session with default chunks433        with client.stream("POST", "/v1/chat/completions", json=payload) as resp:434            complete_lines = read_sse_lines(resp)435 436        # Count total chunks produced437        total_chunks = sum(1 for line in complete_lines438                          if line.startswith("data: ") and '"content"' in line)439 440        if total_chunks == 0:441            pytest.skip("No content chunks produced by default engine")442 443        # Try to resume from early index444        early_resume_index = min(1, total_chunks - 1)  # Resume after first chunk445        headers = {"Last-Event-ID": f"{session_id}:{early_resume_index}"}446 447        with client.stream("POST", "/v1/chat/completions", headers=headers, json=payload) as resp:448            resume_lines = read_sse_lines(resp)449 450        resume_chunks = []451        for line in resume_lines:452            if line.startswith("data: "):453                try:454                    data = json.loads(line[6:])455                    if "choices" in data and data["choices"][0].get("delta", {}).get("content"):456                        resume_chunks.append(data["choices"][0]["delta"]["content"])457                except (json.JSONDecodeError, KeyError):458                    continue459 460        # Should get chunks from resume point onwards, or [DONE] if buffer overflowed461        if resume_chunks:462            # Verify we got some chunks463            assert len(resume_chunks) > 0, "Should get some chunks on resume"464        else:465            # If no chunks, should at least get [DONE]466            assert any("[DONE]" in line for line in resume_lines), "Should get [DONE] even with buffer overflow"467 468 469def test_stream_resume_concurrent_sessions_isolation():470    """Test that resume works correctly with multiple concurrent sessions"""471    with patched_engine():472        client = get_client()473 474        # Create multiple concurrent sessions with different session IDs475        session_ids = ["session_A", "session_B", "session_C"]476        sessions_data = {}477 478        for sid in session_ids:479            payload = {480                "session_id": sid,481                "stream": True,482                "messages": [{"role": "user", "content": f"test {sid}"}],483            }484 485            # Complete session486            with client.stream("POST", "/v1/chat/completions", json=payload) as resp:487                lines = read_sse_lines(resp)488 489            chunks = []490            for line in lines:491                if line.startswith("data: "):492                    try:493                        data = json.loads(line[6:])494                        if "choices" in data and data["choices"][0].get("delta", {}).get("content"):495                            chunks.append(data["choices"][0]["delta"]["content"])496                    except (json.JSONDecodeError, KeyError):497                        continue498 499            sessions_data[sid] = chunks500 501        # Skip if no chunks received502        if all(len(chunks) == 0 for chunks in sessions_data.values()):503            pytest.skip("No content chunks received for any session")504 505        # Test resume for each session independently506        for sid in session_ids:507            if len(sessions_data[sid]) < 2:508                continue  # Skip sessions with too few chunks509 510            resume_point = 0  # Resume from beginning511            headers = {"Last-Event-ID": f"{sid}:{resume_point}"}512            payload = {513                "session_id": sid,514                "stream": True,515                "messages": [{"role": "user", "content": f"test {sid}"}],516            }517 518            with client.stream("POST", "/v1/chat/completions", headers=headers, json=payload) as resp:519                resume_lines = read_sse_lines(resp)520 521            resume_chunks = []522            for line in resume_lines:523                if line.startswith("data: "):524                    try:525                        data = json.loads(line[6:])526                        if "choices" in data and data["choices"][0].get("delta", {}).get("content"):527                            resume_chunks.append(data["choices"][0]["delta"]["content"])528                    except (json.JSONDecodeError, KeyError):529                        continue530 531            # Should get all chunks again (resume from 0)532            assert resume_chunks == sessions_data[sid], (533                f"Session {sid} resume failed: {resume_chunks} != {sessions_data[sid]}"534            )535 536 537def test_stream_resume_with_sqlite_persistence():538    """Test resume works correctly with SQLite persistence enabled"""539    # This test requires setting up SQLite persistence540    original_persist = main.PERSIST_SESSIONS541    original_db_path = main.SESSIONS_DB_PATH542 543    try:544        # Enable persistence for this test545        main.PERSIST_SESSIONS = True546        main.SESSIONS_DB_PATH = ":memory:"  # Use in-memory SQLite for test547 548        # Reinitialize the SQLite store549        main._DB_STORE = main._SQLiteStore(main.SESSIONS_DB_PATH)550 551        with patched_engine():552            client = get_client()553            session_id = "persistent-test"554            payload = {555                "session_id": session_id,556                "stream": True,557                "messages": [{"role": "user", "content": "test persistence"}],558            }559 560            # Complete session to populate SQLite561            with client.stream("POST", "/v1/chat/completions", json=payload) as resp:562                complete_lines = read_sse_lines(resp)563 564            # Extract chunks from complete session565            complete_chunks = []566            for line in complete_lines:567                if line.startswith("data: "):568                    try:569                        data = json.loads(line[6:])570                        if "choices" in data and data["choices"][0].get("delta", {}).get("content"):571                            complete_chunks.append(data["choices"][0]["delta"]["content"])572                    except (json.JSONDecodeError, KeyError):573                        continue574 575            if len(complete_chunks) == 0:576                pytest.skip("No content chunks received for persistence test")577 578            # Verify session was persisted579            assert main._DB_STORE.session_meta(session_id)[0] == True  # Should be marked finished580 581            # Test resume from SQLite582            resume_point = min(1, len(complete_chunks) - 1)583            headers = {"Last-Event-ID": f"{session_id}:{resume_point}"}584 585            with client.stream("POST", "/v1/chat/completions", headers=headers, json=payload) as resp:586                resume_lines = read_sse_lines(resp)587 588            resume_chunks = []589            for line in resume_lines:590                if line.startswith("data: "):591                    try:592                        data = json.loads(line[6:])593                        if "choices" in data and data["choices"][0].get("delta", {}).get("content"):594                            resume_chunks.append(data["choices"][0]["delta"]["content"])595                    except (json.JSONDecodeError, KeyError):596                        continue597 598            expected_chunks = complete_chunks[resume_point:]599            assert resume_chunks == expected_chunks, (600                f"SQLite resume failed: {resume_chunks} != {expected_chunks}"601            )602 603    finally:604        # Restore original settings605        main.PERSIST_SESSIONS = original_persist606        main.SESSIONS_DB_PATH = original_db_path607        main._DB_STORE = main._SQLiteStore(main.SESSIONS_DB_PATH) if main.PERSIST_SESSIONS else None608 609 610def test_stream_resume_data_integrity_with_unicode():611    """Test resume preserves Unicode characters correctly"""612    with patched_engine():613        client = get_client()614        session_id = "unicode-test"615        payload = {616            "session_id": session_id,617            "stream": True,618            "messages": [{"role": "user", "content": "test unicode"}],619        }620 621        # Complete session622        with client.stream("POST", "/v1/chat/completions", json=payload) as resp:623            complete_lines = read_sse_lines(resp)624 625        complete_chunks = []626        for line in complete_lines:627            if line.startswith("data: "):628                try:629                    data = json.loads(line[6:])630                    if "choices" in data and data["choices"][0].get("delta", {}).get("content"):631                        complete_chunks.append(data["choices"][0]["delta"]["content"])632                except (json.JSONDecodeError, KeyError):633                    continue634 635        if len(complete_chunks) == 0:636            pytest.skip("No content chunks received for unicode test")637 638        # Test resume from middle639        resume_point = min(1, len(complete_chunks) - 1)640        headers = {"Last-Event-ID": f"{session_id}:{resume_point}"}641 642        with client.stream("POST", "/v1/chat/completions", headers=headers, json=payload) as resp:643            resume_lines = read_sse_lines(resp)644 645        resume_chunks = []646        for line in resume_lines:647            if line.startswith("data: "):648                try:649                    data = json.loads(line[6:])650                    if "choices" in data and data["choices"][0].get("delta", {}).get("content"):651                        resume_chunks.append(data["choices"][0]["delta"]["content"])652                except (json.JSONDecodeError, KeyError):653                    continue654 655        expected_resume_chunks = complete_chunks[resume_point:]656        assert resume_chunks == expected_resume_chunks657 658        # Verify Unicode integrity (basic check)659        for actual, expected in zip(resume_chunks, expected_resume_chunks):660            assert actual == expected, f"Content mismatch: '{actual}' != '{expected}'"661 662def test_ktp_ocr_success():663    # Mock RapidOCR to return test text lines that should parse to expected KTP data664    test_ocr_texts = [665        "NIK : 1234567890123456",666        "Nama : JOHN DOE",667        "Tempat/Tgl Lahir : JAKARTA, 01-01-1990",668        "Jenis Kelamin : LAKI-LAKI",669        "Alamat : JL. JEND. SUDIRMAN KAV. 52-53",670        "RT/RW : 001/001",671        "Kel/Desa : SENAYAN",672        "Kecamatan : KEBAYORAN BARU",673        "Agama : ISLAM",674        "Status Perkawinan : KAWIN",675        "Pekerjaan : PEGAWAI SWASTA",676        "Kewarganegaraan : WNI",677        "Berlaku Hingga : SEUMUR HIDUP"678    ]679 680    # Mock the OCR result format: [[(bbox, text, confidence), ...]]681    mock_ocr_result = [[(None, text, 0.9) for text in test_ocr_texts]]682 683    # Patch get_ocr_engine to return a mock OCR engine684    original_get_ocr_engine = main.get_ocr_engine685    mock_engine = lambda img: mock_ocr_result686    main.get_ocr_engine = lambda: mock_engine687 688    try:689        client = get_client()690        with open("image.jpg", "rb") as f:691            files = {"image": ("image.jpg", f, "image/jpeg")}692            r = client.post("/ktp-ocr/", files=files)693 694        assert r.status_code == 200695        body = r.json()696        assert body["nik"] == "1234567890123456"697        assert body["nama"] == "John Doe"698        assert body["tempat_lahir"] == "Jakarta"699        assert body["tgl_lahir"] == "01-01-1990"700        assert body["jenis_kelamin"] == "LAKI-LAKI"701        assert body["alamat"]["name"] == "JL. JEND. SUDIRMAN KAV. 52-53"702        assert body["alamat"]["rt_rw"] == "001/001"703        assert body["alamat"]["kel_desa"] == "Senayan"704        assert body["alamat"]["kecamatan"] == "Kebayoran Baru"705        assert body["agama"] == "Islam"706        assert body["status_perkawinan"] == "Kawin"707        assert body["pekerjaan"] == "Pegawai Swasta"708        assert body["kewarganegaraan"] == "Wni"709        assert body["berlaku_hingga"] == "Seumur Hidup"710    finally:711        # Restore original function712        main.get_ocr_engine = original_get_ocr_engine