KillerKing93/Transformers-InferenceServer-OpenAPI
0
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