Download tests/prefix_cache_api_test.py from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 7 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tests/prefix_cache_api_test.py
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/tests/prefix_cache_api_test.py
-
curl -L -o prefix_cache_api_test.py https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tests/prefix_cache_api_test.py
7 kB
| #!/usr/bin/env python3 | |
| """Compare cached/uncached responses, branching, isolation, streaming and cancellation.""" | |
| import argparse | |
| import concurrent.futures | |
| import json | |
| from pathlib import Path | |
| import time | |
| import urllib.error | |
| import urllib.request | |
| p = argparse.ArgumentParser(description=__doc__) | |
| p.add_argument("--url", required=True) | |
| p.add_argument("--output", type=Path, required=True) | |
| a = p.parse_args() | |
| client = urllib.request.build_opener(urllib.request.ProxyHandler({})) | |
| def call(path, body=None): | |
| req = urllib.request.Request(a.url.rstrip("/") + path, | |
| None if body is None else json.dumps(body).encode(), {"Content-Type": "application/json"}) | |
| with client.open(req, timeout=300) as response: | |
| return json.load(response) | |
| def messages(text): | |
| return [{"role": "user", "content": text}] | |
| def body(msgs, cache=True, user="test", count=16): | |
| return {"model": "mindnano-ling3-tiny", "messages": msgs, "max_tokens": count, | |
| "temperature": 0, "cache_prompt": cache, "user": user, "reuse_generated_state": False} | |
| def infer(msgs, cache=True, user="test", count=16): | |
| return call("/v1/chat/completions", body(msgs, cache, user, count)) | |
| def equivalent(reference, actual): | |
| assert actual["choices"] == reference["choices"], "cached reply differs from uncached" | |
| for key in ("prompt_tokens", "completion_tokens", "total_tokens"): | |
| assert actual["usage"][key] == reference["usage"][key], key | |
| m = actual["mindnano_metrics"] | |
| assert m["cached_tokens"] + m["prompt_evaluated_tokens"] == actual["usage"]["prompt_tokens"] | |
| assert actual["usage"]["prompt_tokens_details"]["cached_tokens"] == m["cached_tokens"] | |
| result = {"health": call("/health"), "cases": []} | |
| def record(name, response, reference=None): | |
| if reference is not None: | |
| equivalent(reference, response) | |
| result["cases"].append({"name": name, "response": response, "reference": reference}) | |
| print(json.dumps({"name": name, "metrics": response["mindnano_metrics"]}), flush=True) | |
| def pair(name, seed, target, expected, user="test", seed_user="test"): | |
| cold = infer(target, False, user) | |
| infer(seed, True, seed_user) | |
| hit = infer(target, True, user) | |
| assert hit["mindnano_metrics"]["cached_tokens"] == expected | |
| record(name, hit, cold) | |
| return hit, cold | |
| try: | |
| call("/v1/cache/clear", {}) | |
| short = messages("请用一句话介绍你自己。") | |
| cold = infer(short, False) | |
| warm = infer(short) | |
| equivalent(cold, warm) | |
| exact = infer(short) | |
| assert exact["mindnano_metrics"]["cached_tokens"] == cold["usage"]["prompt_tokens"] | |
| assert exact["mindnano_metrics"]["prompt_evaluated_tokens"] == 0 | |
| record("exact_short", exact, cold) | |
| seed = messages("1" * 1003) # Exactly 1024 tokens including the chat template. | |
| seed_reply = infer(seed) | |
| assert seed_reply["usage"]["prompt_tokens"] == 1024 | |
| follow = seed + [{"role": "assistant", "content": seed_reply["choices"][0]["message"]["content"]}, | |
| {"role": "user", "content": "请用中文简短说明刚才的内容。"}] | |
| follow_hit, follow_cold = pair("conversation_1024", seed, follow, 1024) | |
| follow_again = infer(follow) | |
| assert follow_again["mindnano_metrics"]["prompt_evaluated_tokens"] == 0 | |
| record("exact_followup", follow_again, follow_cold) | |
| partial = messages("1" * 236) | |
| pair("partial_block", partial, partial + [ | |
| {"role": "assistant", "content": "收到"}, {"role": "user", "content": "请继续说明。"}], 256) | |
| pair("tail_edit_after_checkpoint", messages("1" * 300), messages("1" * 280 + "2" * 20), 256) | |
| pair("edit_before_checkpoint", messages("1" * 300), messages("2" + "1" * 299), 0) | |
| pair("shorter_request", messages("1" * 300), short, 0) | |
| pair("user_isolation", seed, seed, 0, user="bob", seed_user="alice") | |
| pair("checkpoint_only_shorter", messages("1" * 236), messages("1" * 235), 0) | |
| # Re-enter the same prefix after an unrelated request: this is a bounded single entry. | |
| cold = infer(seed, False) | |
| infer(seed) | |
| infer(short) | |
| evicted = infer(seed) | |
| assert evicted["mindnano_metrics"]["cached_tokens"] == 0 | |
| record("single_entry_eviction", evicted, cold) | |
| # Streaming must still use saved logits and restore pre-generation recurrent state. | |
| request_body = body(seed) | |
| request_body.update(stream=True, stream_options={"include_usage": True}) | |
| request = urllib.request.Request(a.url.rstrip("/") + "/v1/chat/completions", | |
| json.dumps(request_body).encode(), {"Content-Type": "application/json"}) | |
| content, metrics, usage, done = "", None, None, False | |
| with client.open(request, timeout=300) as response: | |
| for raw in response: | |
| line = raw.decode().strip() | |
| if not line.startswith("data: "): | |
| continue | |
| if line == "data: [DONE]": | |
| done = True | |
| break | |
| event = json.loads(line[6:]) | |
| assert "error" not in event, event | |
| for choice in event["choices"]: | |
| content += choice["delta"].get("content", "") | |
| metrics = event.get("mindnano_metrics", metrics) | |
| usage = event.get("usage") or usage | |
| assert done and metrics["cached_tokens"] == 1024 and usage | |
| assert content == cold["choices"][0]["message"]["content"] | |
| result["streaming"] = {"passed": True, "metrics": metrics} | |
| # Cancel a request with a valid reused prefix; all cached state must be discarded. | |
| long = seed + [{"role": "assistant", "content": "收到"}, | |
| {"role": "user", "content": "1" * 2000}] | |
| with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool: | |
| future = pool.submit(infer, long) | |
| for _ in range(200): | |
| if call("/health")["active"]: | |
| break | |
| time.sleep(0.02) | |
| else: | |
| raise AssertionError("long request did not start") | |
| for path, data in (("/v1/cache/clear", {}), ("/v1/chat/completions", body(short))): | |
| try: | |
| call(path, data) | |
| raise AssertionError("concurrent mutation/inference accepted") | |
| except urllib.error.HTTPError as exc: | |
| assert exc.code == 429 | |
| assert call("/v1/cancel", {})["cancel_requested"] | |
| try: | |
| future.result(timeout=60) | |
| raise AssertionError("canceled request returned a completed reply") | |
| except urllib.error.HTTPError as exc: | |
| assert exc.code == 409 | |
| after_cancel = infer(seed) | |
| assert after_cancel["mindnano_metrics"]["cached_tokens"] == 0 | |
| record("cancel_clears_cache", after_cancel, cold) | |
| call("/v1/cache/clear", {}) | |
| after_clear = infer(seed) | |
| assert after_clear["mindnano_metrics"]["cached_tokens"] == 0 | |
| record("explicit_clear", after_clear, cold) | |
| result["passed"] = True | |
| finally: | |
| a.output.parent.mkdir(parents=True, exist_ok=True) | |
| a.output.write_text(json.dumps(result, ensure_ascii=False, indent=2) + "\n") | |