File size: 5,262 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
"""Default QA cache, resident session and lifecycle checks on a dedicated engine."""
import argparse
import json
from pathlib import Path
import urllib.request


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--base", required=True)
    parser.add_argument("--output", type=Path, required=True)
    args = parser.parse_args()
    client = urllib.request.build_opener(urllib.request.ProxyHandler({}))

    def call(path, body=None):
        with 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) as response:
            return json.load(response)

    def infer(messages, session="qa", count=16, **extra):
        body = {"model": "mindnano-ling3-tiny", "messages": messages,
                "user": "qa-default", "max_tokens": count, "temperature": 0}
        if session is not None:
            body["session_id"] = session
        return call("/v1/chat/completions", dict(body, **extra))

    def clear():
        call("/v1/cache/clear", {})

    def follow(messages, reply, question="请简短说明刚才的内容。"):
        return messages + [reply["choices"][0]["message"], {"role": "user", "content": question}]

    result = {"passed": False, "chains": []}
    assert call("/health")["capabilities"]["qa_cache_default"]
    assert not call("/v1/generation/status")["active"]
    clear()
    try:
        # No reuse_generated_state parameter anywhere in these requests.
        for session in (None, "qa"):
            clear()
            messages = [{"role": "user", "content": "1" * 107}]
            replies = []
            for turn in range(4):
                response = infer(messages, session, count=32 if not turn else 16)
                m = response["mindnano_metrics"]
                assert m["reuse_generated_state"]
                assert m["cached_tokens"] + m["prompt_evaluated_tokens"] == response["usage"]["prompt_tokens"]
                if turn == 1:
                    assert m["cache_status"] == "generated_prefix_hit"
                    assert m["generated_cached_tokens"] == 31
                elif turn > 1:
                    assert m["cached_tokens"] >= 128
                replies.append(response)
                messages = follow(messages, response, "请简短说明刚才的内容。" if not turn else "请再用一句话总结。")
            result["chains"].append(replies)
        for anonymous, named in zip(*result["chains"]):
            assert anonymous["choices"] == named["choices"], "resident named state differs from uninterrupted anonymous state"

        # Short questions and one-token replies still reuse the exact Q prefix.
        clear()
        short = [{"role": "user", "content": "请只回答数字1。"}]
        first = infer(short, count=1)
        second = infer(follow(short, first))
        assert first["usage"]["prompt_tokens"] < 128
        assert second["mindnano_metrics"]["cache_status"] == "exact_prefix_hit"
        assert second["mindnano_metrics"]["cached_tokens"] == first["usage"]["prompt_tokens"]
        result["short_question"] = second

        # Inputs larger than the history budget keep the active named QA state.
        clear()
        long_input = [{"role": "user", "content": "1" * 1003}]
        first = infer(long_input)
        second = infer(follow(long_input, first))
        assert second["mindnano_metrics"]["cached_tokens"] == 1039
        assert second["mindnano_metrics"]["generated_cached_tokens"] == 15
        status = call("/v1/cache/status")
        assert status["resident_named"] and status["resident_cache_bytes"] > 0
        assert status["snapshot_bytes"] <= status["budget_bytes"]
        if status["budget_bytes"] < 64 * 1024 * 1024:
            assert status["sessions"] == 0
            assert second["mindnano_metrics"]["session_cache_skipped_budget"]
        result["long_resident"] = {"response": second, "cache": status}

        # Scoped clear must invalidate resident state, not only historical copies.
        call("/v1/cache/clear", {"user": "qa-default", "session_id": "qa"})
        assert call("/v1/cache/status")["resident_cache_bytes"] == 0
        cold = infer(short)
        assert cold["mindnano_metrics"]["cached_tokens"] == 0
        infer(short, cache_prompt=False)
        assert call("/v1/cache/status")["resident_cache_bytes"] == 0
        assert infer(short)["mindnano_metrics"]["cached_tokens"] == 0

        if status["budget_bytes"] >= 256 * 1024 * 1024:
            clear()
            a = infer(short, "a")
            infer([{"role": "user", "content": "介绍冬天。"}], "b")
            call("/v1/cache/fork", {"user": "qa-default", "source_session_id": "a", "target_session_id": "b"})
            b = infer(short, "b")
            assert b["choices"] == a["choices"] and b["mindnano_metrics"]["cached_tokens"] > 0
        result["passed"] = True
    finally:
        clear()
        args.output.write_text(json.dumps(result, ensure_ascii=False, indent=2) + "\n")
        print(json.dumps({"qa_cache": "PASS" if result["passed"] else "FAIL", "output": str(args.output)}), flush=True)


if __name__ == "__main__":
    main()