File size: 7,463 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
"""Real inference regression. Run only against a dedicated test engine."""
import argparse
import json
import time
import urllib.error
import urllib.request


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--base", required=True)
    parser.add_argument("--budget-test", action="store_true",
                        help="Use a dedicated engine with 0 or 48 MiB session budget")
    args = parser.parse_args()
    client = urllib.request.build_opener(urllib.request.ProxyHandler({}))

    def open_request(path, body=None):
        return client.open(urllib.request.Request(args.base + path,
            data=None if body is None else json.dumps(body).encode(),
            headers={"Content-Type": "application/json"}), timeout=180)

    def call(path, body=None):
        with open_request(path, body) as response:
            return json.load(response)

    def body(session, text="23乘17等于多少?", user="state-test"):
        return {"model": "mindnano-ling3-tiny", "messages": [{"role": "user", "content": text}],
                "session_id": session, "user": user, "temperature": 0, "max_tokens": 8, "reuse_generated_state": False}

    def infer(session, text="23乘17等于多少?", user="state-test", **extra):
        return call("/v1/chat/completions", dict(body(session, text, user), **extra))

    def equivalent(a, b):
        assert a["choices"] == b["choices"], "restored generation changed"
        assert a["usage"]["completion_tokens"] == b["usage"]["completion_tokens"]

    def hit(result):
        return result["mindnano_metrics"]["cached_tokens"]

    def rejected(path, payload, code):
        try:
            call(path, payload)
        except urllib.error.HTTPError as error:
            assert error.code == code
            return
        raise AssertionError("invalid cache management request accepted")

    assert not call("/v1/generation/status")["active"]
    assert call("/health")["capabilities"]["session_cache"]
    call("/v1/cache/clear", {})
    try:
        if args.budget_test:
            status = call("/v1/cache/status")
            assert status["budget_bytes"] in (0, 48 * 1024 * 1024)
            first = infer("a")
            infer("b", "用一句话介绍春天。")
            again = infer("a"); equivalent(first, again); assert hit(again) == 0
            if status["budget_bytes"]:
                assert again["mindnano_metrics"]["session_cache_stored"]
                assert call("/v1/cache/status")["evictions"] >= 2
            else:
                assert again["mindnano_metrics"]["session_cache_skipped_budget"]
            # A request exceeding the cache budget still computes all its input.
            large = infer("a", "1" * 300)
            assert large["usage"]["prompt_tokens"] == 321
            assert large["mindnano_metrics"]["prompt_evaluated_tokens"] == 321
            assert large["mindnano_metrics"]["session_cache_skipped_budget"]
            assert not large["mindnano_metrics"]["session_cache_stored"]
            assert hit(infer("a")) == 0
            if status["budget_bytes"]:
                for enabled in (False, True):
                    answers = []
                    for instruction in ("on", "off"):
                        answers.append(call("/v1/chat/completions", dict(body("thinking"),
                            enable_thinking=enabled, max_tokens=32, messages=[
                                {"role": "system", "content": "detailed thinking " + instruction},
                                {"role": "user", "content": "23乘17等于多少?"}])))
                    equivalent(*answers)
                    assert hit(answers[1]) == answers[1]["usage"]["prompt_tokens"]
                print(json.dumps({"thinking_switch_overrides_conflicting_system": "PASS"}), flush=True)
            print(json.dumps({"cache_budget": "PASS", "no_input_truncation": "PASS",
                              "cache": call("/v1/cache/status")}), flush=True)
            return
        first = infer("a")
        before = call("/v1/cache/status")
        rejected("/v1/cache/clear", {"sessoin_id": "a"}, 400)
        rejected("/v1/cache/clear", {"user": "state-test"}, 400)
        rejected("/v1/cache/fork", {"source_session_id": "absent", "target_session_id": "new"}, 404)
        assert call("/v1/cache/status") == before
        long_user = "u" * 300
        infer("long-user", user=long_user)
        call("/v1/cache/fork", {"user": long_user, "source_session_id": "long-user", "target_session_id": "copy"})
        call("/v1/cache/clear", {"user": long_user, "session_id": "long-user"})
        call("/v1/cache/clear", {"user": long_user, "session_id": "copy"})
        infer("b", "用一句话介绍春天。")
        again = infer("a")
        equivalent(first, again)
        assert hit(again) == first["usage"]["prompt_tokens"]
        call("/v1/cache/fork", {"user": "state-test", "source_session_id": "a", "target_session_id": "fork"})
        fork = infer("fork"); equivalent(first, fork); assert hit(fork) > 0
        infer("fork", "你好,请介绍秋天。")
        equivalent(first, infer("a"))
        assert hit(infer("a", user="another-user")) == 0
        call("/v1/cache/clear", {"session_id": "b", "user": "state-test"})
        assert hit(infer("a")) > 0
        assert hit(infer("b", "用一句话介绍春天。")) == 0
        # Reuse an aligned prefix across another session; compare to cold run.
        text = "1" * 140
        seeded = infer("prefix", text)
        extended = text + "2" * 10
        cold = infer("cold", extended, cache_prompt=False)
        warm = infer("prefix", extended)
        equivalent(cold, warm); assert hit(warm) >= 128
        equivalent(first, infer("a"))  # Refresh its residency before cancellation.
        # Cancel an ACK-parked update. Its previous successful prompt survives.
        request_id = None
        try:
            with open_request("/v1/chat/completions", dict(body("a", "讲一个长故事。"),
                              stream=True, flow_control="ack", max_tokens=128)) as response:
                for line in response:
                    if not line.startswith(b"data:"):
                        continue
                    event = json.loads(line[5:])
                    if event.get("mindnano_flow"):
                        request_id = event["id"]
                        break
                assert request_id
                call("/v1/cancel", {"request_id": request_id})
                for line in response:
                    if line.strip() == b"data: [DONE]":
                        break
        finally:
            if request_id:
                call("/v1/cancel", {"request_id": request_id})
        deadline = time.monotonic() + 10
        while call("/v1/generation/status")["active"]:
            assert time.monotonic() < deadline
            time.sleep(.05)
        resumed = infer("a"); equivalent(first, resumed); assert hit(resumed) > 0
        status = call("/v1/cache/status")
        assert status["snapshot_bytes"] <= status["budget_bytes"]
        print(json.dumps({"session_isolation": "PASS", "fork": "PASS", "prefix": "PASS",
                          "cancel_preserves_committed_cache": "PASS", "cache": status,
                          "exact_hit_ttft_ms": again["mindnano_metrics"]["ttft_ms"]}), flush=True)
    finally:
        call("/v1/cache/clear", {})


if __name__ == "__main__":
    main()