Ling-3.0-tiny-RKNN / tests /prefix_cache_api_test.py
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
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")