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")