File size: 6,998 Bytes
3fd1a35 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 | #!/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")
|