@modular-prompt/driver 0.16.0 → 0.17.1

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (202) hide show
  1. package/README.md +98 -10
  2. package/dist/cache-controller.d.ts +4 -0
  3. package/dist/cache-controller.d.ts.map +1 -1
  4. package/dist/driver-registry/config-based-factory.d.ts +3 -0
  5. package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
  6. package/dist/driver-registry/config-based-factory.js +8 -1
  7. package/dist/driver-registry/config-based-factory.js.map +1 -1
  8. package/dist/driver-registry/factory-helper.d.ts.map +1 -1
  9. package/dist/driver-registry/factory-helper.js +9 -2
  10. package/dist/driver-registry/factory-helper.js.map +1 -1
  11. package/dist/driver-registry/index.d.ts +1 -1
  12. package/dist/driver-registry/index.d.ts.map +1 -1
  13. package/dist/driver-registry/types.d.ts +15 -1
  14. package/dist/driver-registry/types.d.ts.map +1 -1
  15. package/dist/formatter/converter.d.ts.map +1 -1
  16. package/dist/formatter/converter.js +31 -2
  17. package/dist/formatter/converter.js.map +1 -1
  18. package/dist/google-genai/google-genai-driver.d.ts +1 -0
  19. package/dist/google-genai/google-genai-driver.d.ts.map +1 -1
  20. package/dist/google-genai/google-genai-driver.js +36 -27
  21. package/dist/google-genai/google-genai-driver.js.map +1 -1
  22. package/dist/index.d.ts +5 -3
  23. package/dist/index.d.ts.map +1 -1
  24. package/dist/index.js +5 -3
  25. package/dist/index.js.map +1 -1
  26. package/dist/local-inference/adapters.d.ts +6 -0
  27. package/dist/local-inference/adapters.d.ts.map +1 -1
  28. package/dist/local-inference/driver.d.ts.map +1 -1
  29. package/dist/local-inference/driver.js +45 -24
  30. package/dist/local-inference/driver.js.map +1 -1
  31. package/dist/local-inference/process-client.d.ts +4 -2
  32. package/dist/local-inference/process-client.d.ts.map +1 -1
  33. package/dist/local-inference/process-client.js +24 -8
  34. package/dist/local-inference/process-client.js.map +1 -1
  35. package/dist/local-inference/process-communication.d.ts +9 -2
  36. package/dist/local-inference/process-communication.d.ts.map +1 -1
  37. package/dist/local-inference/process-communication.js +37 -5
  38. package/dist/local-inference/process-communication.js.map +1 -1
  39. package/dist/local-inference/protocol.d.ts +4 -0
  40. package/dist/local-inference/protocol.d.ts.map +1 -1
  41. package/dist/local-inference/request-queue.d.ts +1 -1
  42. package/dist/local-inference/request-queue.d.ts.map +1 -1
  43. package/dist/local-inference/request-queue.js +26 -7
  44. package/dist/local-inference/request-queue.js.map +1 -1
  45. package/dist/local-inference/stream-utils.d.ts +6 -0
  46. package/dist/local-inference/stream-utils.d.ts.map +1 -1
  47. package/dist/local-inference/stream-utils.js.map +1 -1
  48. package/dist/mlx-ml/mlx-cache-controller.d.ts +9 -0
  49. package/dist/mlx-ml/mlx-cache-controller.d.ts.map +1 -1
  50. package/dist/mlx-ml/mlx-cache-controller.js +158 -33
  51. package/dist/mlx-ml/mlx-cache-controller.js.map +1 -1
  52. package/dist/mlx-ml/mlx-cache-support.d.ts +2 -2
  53. package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -1
  54. package/dist/mlx-ml/mlx-cache-support.js +8 -3
  55. package/dist/mlx-ml/mlx-cache-support.js.map +1 -1
  56. package/dist/mlx-ml/mlx-driver.d.ts +0 -1
  57. package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
  58. package/dist/mlx-ml/mlx-driver.js +1 -8
  59. package/dist/mlx-ml/mlx-driver.js.map +1 -1
  60. package/dist/mlx-ml/process/index.d.ts +1 -1
  61. package/dist/mlx-ml/process/index.d.ts.map +1 -1
  62. package/dist/mlx-ml/process/index.js +2 -2
  63. package/dist/mlx-ml/process/index.js.map +1 -1
  64. package/dist/models-config/index.d.ts +1 -1
  65. package/dist/models-config/index.d.ts.map +1 -1
  66. package/dist/models-config/index.js +1 -1
  67. package/dist/models-config/index.js.map +1 -1
  68. package/dist/models-config/resolve.d.ts +9 -1
  69. package/dist/models-config/resolve.d.ts.map +1 -1
  70. package/dist/models-config/resolve.js +94 -2
  71. package/dist/models-config/resolve.js.map +1 -1
  72. package/dist/models-config/types.d.ts +3 -1
  73. package/dist/models-config/types.d.ts.map +1 -1
  74. package/dist/pytorch/process/index.d.ts +4 -2
  75. package/dist/pytorch/process/index.d.ts.map +1 -1
  76. package/dist/pytorch/process/index.js +24 -7
  77. package/dist/pytorch/process/index.js.map +1 -1
  78. package/dist/pytorch/pytorch-cache-controller.d.ts +84 -0
  79. package/dist/pytorch/pytorch-cache-controller.d.ts.map +1 -0
  80. package/dist/pytorch/pytorch-cache-controller.js +742 -0
  81. package/dist/pytorch/pytorch-cache-controller.js.map +1 -0
  82. package/dist/pytorch/pytorch-cache-support.d.ts +23 -0
  83. package/dist/pytorch/pytorch-cache-support.d.ts.map +1 -0
  84. package/dist/pytorch/pytorch-cache-support.js +47 -0
  85. package/dist/pytorch/pytorch-cache-support.js.map +1 -0
  86. package/dist/pytorch/pytorch-driver.d.ts +8 -1
  87. package/dist/pytorch/pytorch-driver.d.ts.map +1 -1
  88. package/dist/pytorch/pytorch-driver.js +40 -0
  89. package/dist/pytorch/pytorch-driver.js.map +1 -1
  90. package/dist/runtime/check.d.ts.map +1 -1
  91. package/dist/runtime/check.js +8 -6
  92. package/dist/runtime/check.js.map +1 -1
  93. package/dist/runtime/index.d.ts +2 -2
  94. package/dist/runtime/index.d.ts.map +1 -1
  95. package/dist/runtime/index.js +2 -2
  96. package/dist/runtime/index.js.map +1 -1
  97. package/dist/runtime/manifest-core.d.mts +1 -0
  98. package/dist/runtime/manifest-core.mjs +1 -0
  99. package/dist/runtime/manifest-core.mjs.map +1 -1
  100. package/dist/runtime/manifest.d.ts +2 -0
  101. package/dist/runtime/manifest.d.ts.map +1 -1
  102. package/dist/runtime/manifest.js.map +1 -1
  103. package/dist/runtime/paths-core.d.mts +15 -1
  104. package/dist/runtime/paths-core.d.mts.map +1 -1
  105. package/dist/runtime/paths-core.mjs +50 -5
  106. package/dist/runtime/paths-core.mjs.map +1 -1
  107. package/dist/runtime/paths.d.ts +2 -2
  108. package/dist/runtime/paths.d.ts.map +1 -1
  109. package/dist/runtime/paths.js +2 -2
  110. package/dist/runtime/paths.js.map +1 -1
  111. package/dist/runtime/pytorch-template-core.d.mts +11 -0
  112. package/dist/runtime/pytorch-template-core.d.mts.map +1 -0
  113. package/dist/runtime/pytorch-template-core.mjs +54 -0
  114. package/dist/runtime/pytorch-template-core.mjs.map +1 -0
  115. package/dist/runtime/setup-commands-core.d.mts +3 -0
  116. package/dist/runtime/setup-commands-core.d.mts.map +1 -1
  117. package/dist/runtime/setup-commands-core.mjs +4 -0
  118. package/dist/runtime/setup-commands-core.mjs.map +1 -1
  119. package/dist/runtime/setup-commands.d.ts +1 -1
  120. package/dist/runtime/setup-commands.d.ts.map +1 -1
  121. package/dist/runtime/setup-commands.js +1 -1
  122. package/dist/runtime/setup-commands.js.map +1 -1
  123. package/dist/vertexai/vertexai-driver.d.ts +6 -0
  124. package/dist/vertexai/vertexai-driver.d.ts.map +1 -1
  125. package/dist/vertexai/vertexai-driver.js +106 -36
  126. package/dist/vertexai/vertexai-driver.js.map +1 -1
  127. package/docs/DRIVER_API.md +455 -0
  128. package/docs/LOCAL_MODEL_SETUP.md +765 -0
  129. package/docs/mlx-api-selection.md +301 -0
  130. package/package.json +10 -6
  131. package/scripts/runtime-cli.bin.test.ts +142 -0
  132. package/scripts/runtime-cli.js +305 -35
  133. package/scripts/runtime-cli.test.ts +163 -0
  134. package/skills/driver-usage/SKILL.md +29 -0
  135. package/src/mlx-ml/python/__main__.py +1 -1
  136. package/src/mlx-ml/python/backends/base.py +88 -18
  137. package/src/mlx-ml/python/backends/mlx_lm.py +28 -3
  138. package/src/mlx-ml/python/backends/mlx_vlm.py +679 -2
  139. package/src/mlx-ml/python/handlers/cache.py +4 -0
  140. package/src/mlx-ml/python/handlers/generate.py +33 -10
  141. package/src/mlx-ml/python/handlers/tokenize.py +1 -4
  142. package/src/mlx-ml/python/pyproject.toml +2 -2
  143. package/src/mlx-ml/python/server.py +2 -0
  144. package/src/mlx-ml/python/uv.lock +12 -12
  145. package/src/pytorch/templates/cpu-minimal/backends/base.py +139 -0
  146. package/src/pytorch/templates/cpu-minimal/backends/transformers_lm.py +1167 -0
  147. package/src/pytorch/{python → templates/cpu-minimal}/handlers/__init__.py +1 -0
  148. package/src/pytorch/templates/cpu-minimal/handlers/cache.py +88 -0
  149. package/src/pytorch/templates/cpu-minimal/handlers/generate.py +157 -0
  150. package/src/pytorch/{python → templates/cpu-minimal}/pyproject.toml +2 -2
  151. package/src/pytorch/{python → templates/cpu-minimal}/server.py +20 -2
  152. package/src/pytorch/templates/cpu-minimal/tests/test_cache_handler.py +284 -0
  153. package/src/pytorch/templates/cpu-minimal/tests/test_capabilities.py +14 -0
  154. package/src/pytorch/templates/cpu-minimal/tests/test_server.py +141 -0
  155. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_errors.py +89 -0
  156. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_lm_cache.py +554 -0
  157. package/src/pytorch/templates/cpu-minimal/utils/__init__.py +0 -0
  158. package/src/pytorch/{python → templates/cpu-minimal}/utils/token_utils.py +2 -2
  159. package/src/pytorch/templates/cpu-minimal/utils/transformers_errors.py +54 -0
  160. package/src/pytorch/{python → templates/cpu-minimal}/uv.lock +149 -109
  161. package/src/pytorch/templates/cuda/__main__.py +19 -0
  162. package/src/pytorch/templates/cuda/backends/__init__.py +3 -0
  163. package/src/pytorch/{python → templates/cuda}/backends/base.py +54 -6
  164. package/src/pytorch/templates/cuda/backends/transformers_lm.py +379 -0
  165. package/src/pytorch/templates/cuda/handlers/__init__.py +7 -0
  166. package/src/pytorch/templates/cuda/handlers/cache.py +93 -0
  167. package/src/pytorch/templates/cuda/handlers/cancel.py +53 -0
  168. package/src/pytorch/templates/cuda/handlers/capabilities.py +6 -0
  169. package/src/pytorch/templates/cuda/handlers/completion.py +15 -0
  170. package/src/pytorch/templates/cuda/handlers/format_test.py +70 -0
  171. package/src/pytorch/templates/cuda/handlers/generate.py +152 -0
  172. package/src/pytorch/templates/cuda/handlers/render.py +40 -0
  173. package/src/pytorch/templates/cuda/handlers/tokenize.py +63 -0
  174. package/src/pytorch/templates/cuda/pyproject.toml +37 -0
  175. package/src/pytorch/templates/cuda/server.py +158 -0
  176. package/src/pytorch/templates/cuda/tests/__init__.py +0 -0
  177. package/src/pytorch/templates/cuda/tests/test_cache_handler.py +207 -0
  178. package/src/pytorch/templates/cuda/tests/test_capabilities.py +14 -0
  179. package/src/pytorch/templates/cuda/tests/test_server.py +145 -0
  180. package/src/pytorch/templates/cuda/tests/test_transformers_errors.py +89 -0
  181. package/src/pytorch/templates/cuda/tests/test_transformers_lm_cache.py +288 -0
  182. package/src/pytorch/templates/cuda/utils/__init__.py +0 -0
  183. package/src/pytorch/templates/cuda/utils/chat_template_constraints.py +164 -0
  184. package/src/pytorch/templates/cuda/utils/prompt_builder.py +54 -0
  185. package/src/pytorch/templates/cuda/utils/template_render.py +80 -0
  186. package/src/pytorch/templates/cuda/utils/token_utils.py +376 -0
  187. package/src/pytorch/templates/cuda/utils/transformers_errors.py +54 -0
  188. package/src/pytorch/templates/cuda/uv.lock +734 -0
  189. package/src/pytorch/python/backends/transformers_lm.py +0 -127
  190. package/src/pytorch/python/handlers/generate.py +0 -68
  191. /package/src/pytorch/{python → templates/cpu-minimal}/__main__.py +0 -0
  192. /package/src/pytorch/{python → templates/cpu-minimal}/backends/__init__.py +0 -0
  193. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/cancel.py +0 -0
  194. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/capabilities.py +0 -0
  195. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/completion.py +0 -0
  196. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/format_test.py +0 -0
  197. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/render.py +0 -0
  198. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/tokenize.py +0 -0
  199. /package/src/pytorch/{python/utils → templates/cpu-minimal/tests}/__init__.py +0 -0
  200. /package/src/pytorch/{python → templates/cpu-minimal}/utils/chat_template_constraints.py +0 -0
  201. /package/src/pytorch/{python → templates/cpu-minimal}/utils/prompt_builder.py +0 -0
  202. /package/src/pytorch/{python → templates/cpu-minimal}/utils/template_render.py +0 -0
@@ -0,0 +1,152 @@
1
+ from __future__ import annotations
2
+
3
+ import json
4
+
5
+ from backends.base import ModelBackend
6
+ from handlers.cancel import poll_cancel
7
+
8
+
9
+ def _stream_to_stdout(
10
+ backend: ModelBackend,
11
+ prompt: str | list[int],
12
+ options: dict,
13
+ images: list | None = None,
14
+ primer: str | None = None,
15
+ prompt_cache=None,
16
+ cache_loaded: bool | None = None,
17
+ cache_read_tokens: int = 0,
18
+ cache_write_tokens: int = 0,
19
+ ) -> None:
20
+ if images:
21
+ raise ValueError("PyTorch LIP backend does not support images in Phase 1")
22
+
23
+ if primer is not None:
24
+ print(primer, end="", flush=True)
25
+
26
+ first_prompt_tokens = None
27
+ first_cache_read_tokens = None
28
+ first_cache_write_tokens = None
29
+ reported_generation_tokens = None
30
+ for response in backend.stream_generate(
31
+ prompt,
32
+ options,
33
+ images,
34
+ prompt_cache=prompt_cache,
35
+ ):
36
+ if poll_cancel():
37
+ break
38
+ response_prompt_tokens = getattr(response, "prompt_tokens", None)
39
+ if first_prompt_tokens is None and response_prompt_tokens is not None:
40
+ first_prompt_tokens = response_prompt_tokens
41
+ if (
42
+ first_cache_read_tokens is None
43
+ and getattr(response, "cache_read_tokens", None) is not None
44
+ ):
45
+ first_cache_read_tokens = response.cache_read_tokens
46
+ if (
47
+ first_cache_write_tokens is None
48
+ and getattr(response, "cache_write_tokens", None) is not None
49
+ ):
50
+ first_cache_write_tokens = response.cache_write_tokens
51
+ response_generation_tokens = getattr(response, "generation_tokens", None)
52
+ if response_generation_tokens is not None:
53
+ reported_generation_tokens = max(
54
+ reported_generation_tokens or 0,
55
+ int(response_generation_tokens),
56
+ )
57
+ print(response.text.replace("\0", "").replace("\x1e", ""), end="", flush=True)
58
+
59
+ meta: dict = {}
60
+ if first_prompt_tokens is not None:
61
+ meta["prompt_tokens"] = first_prompt_tokens
62
+ if reported_generation_tokens is not None:
63
+ meta["generation_tokens"] = max(0, reported_generation_tokens)
64
+ if first_cache_read_tokens is not None:
65
+ meta["cache_read_tokens"] = first_cache_read_tokens
66
+ if first_cache_write_tokens is not None:
67
+ meta["cache_write_tokens"] = first_cache_write_tokens
68
+ if cache_read_tokens > 0 and "cache_read_tokens" not in meta:
69
+ meta["cache_read_tokens"] = cache_read_tokens
70
+ if cache_write_tokens > 0 and "cache_write_tokens" not in meta:
71
+ meta["cache_write_tokens"] = cache_write_tokens
72
+ if cache_loaded is not None:
73
+ meta["cache_loaded"] = cache_loaded
74
+
75
+ if meta:
76
+ print(f"\x1e__META__:{json.dumps(meta)}", end="\0", flush=True)
77
+ else:
78
+ print("", end="\0", flush=True)
79
+
80
+
81
+ def handle_generate(
82
+ backend: ModelBackend,
83
+ prompt: str | list[int],
84
+ options: dict | None = None,
85
+ images: list | None = None,
86
+ max_image_size: int = 768,
87
+ primer: str | None = None,
88
+ cache_path: str | None = None,
89
+ cache_trim_tokens: int | None = None,
90
+ ) -> None:
91
+ """LIP generate: 整形済み prompt のストリーム推論"""
92
+ if cache_trim_tokens is not None:
93
+ raise ValueError(
94
+ "PyTorch LIP backend does not support cache trimming in Phase 1"
95
+ )
96
+
97
+ if options is None:
98
+ options = {}
99
+
100
+ final_options = dict(options)
101
+ final_options.pop("trust_remote_code", None)
102
+
103
+ prompt_cache = None
104
+ cache_loaded = None
105
+ cache_read_tokens = 0
106
+ cache_write_tokens = 0
107
+ if cache_path:
108
+ if images:
109
+ cache_loaded = False
110
+ else:
111
+ prompt_cache = backend.load_cache_from_file(
112
+ cache_path,
113
+ images=images,
114
+ max_image_size=max_image_size,
115
+ prompt=prompt,
116
+ )
117
+ cache_loaded = prompt_cache is not None
118
+
119
+ if prompt_cache is not None:
120
+ cache_read_tokens = backend.get_cache_offset(prompt_cache)
121
+ if cache_read_tokens <= 0:
122
+ prompt_cache = None
123
+ cache_loaded = False
124
+ elif isinstance(prompt, (str, list)):
125
+ full_tokens = (
126
+ backend.tokenize_prompt(prompt)
127
+ if isinstance(prompt, str)
128
+ else [int(token_id) for token_id in prompt]
129
+ )
130
+ if cache_read_tokens < len(full_tokens):
131
+ prompt = full_tokens[cache_read_tokens:]
132
+ else:
133
+ # A cache covering the complete prompt cannot be passed
134
+ # with an empty input_ids tensor. The safe fallback is a
135
+ # cold generation for this Phase 1 handler.
136
+ prompt_cache = None
137
+ cache_loaded = False
138
+ cache_read_tokens = 0
139
+ if prompt_cache is not None:
140
+ cache_write_tokens = backend.consume_cache_write_tokens(cache_path)
141
+
142
+ _stream_to_stdout(
143
+ backend,
144
+ prompt,
145
+ final_options,
146
+ images=images,
147
+ primer=primer,
148
+ prompt_cache=prompt_cache,
149
+ cache_loaded=cache_loaded,
150
+ cache_read_tokens=cache_read_tokens,
151
+ cache_write_tokens=cache_write_tokens,
152
+ )
@@ -0,0 +1,40 @@
1
+ from __future__ import annotations
2
+
3
+ import json
4
+
5
+ from backends.base import ModelBackend
6
+ from utils.template_render import apply_chat_template_prompt
7
+
8
+
9
+ def handle_render(
10
+ backend: ModelBackend,
11
+ messages: list,
12
+ options: dict | None = None,
13
+ tools: list | None = None,
14
+ reasoning_effort: str | None = None,
15
+ ) -> None:
16
+ """LIP render: apply_chat_template のみ(推論しない)"""
17
+ if options is None:
18
+ options = {}
19
+
20
+ result: dict = {
21
+ "formatted_prompt": None,
22
+ "error": None,
23
+ }
24
+
25
+ try:
26
+ trust_remote_code = options.get("trust_remote_code")
27
+ primer = options.get("primer")
28
+ prompt = apply_chat_template_prompt(
29
+ backend,
30
+ messages,
31
+ primer=primer,
32
+ tools=tools,
33
+ reasoning_effort=reasoning_effort,
34
+ trust_remote_code=trust_remote_code,
35
+ )
36
+ result["formatted_prompt"] = prompt
37
+ except Exception as e:
38
+ result["error"] = str(e)
39
+
40
+ print(json.dumps(result), end="\0", flush=True)
@@ -0,0 +1,63 @@
1
+ import json
2
+
3
+ from backends.base import ModelBackend
4
+ from utils.prompt_builder import generate_merged_prompt, supports_chat_template
5
+
6
+
7
+ def handle_tokenize(
8
+ backend: ModelBackend,
9
+ capabilities: dict,
10
+ messages: list,
11
+ tools: list | None = None,
12
+ reasoning_effort: str | None = None,
13
+ ) -> None:
14
+ """メッセージをchat template適用後にトークン化して返す"""
15
+ tokenizer = backend.get_tokenizer()
16
+
17
+ result = {
18
+ "token_ids": None,
19
+ "token_count": 0,
20
+ "error": None,
21
+ }
22
+
23
+ try:
24
+ # apply_chat_templateのfallbackパターン (chat.py L165-188 と同じ)
25
+ # add_generation_prompt=False で、アシスタントの開始トークンは含めない
26
+ extra_kwargs = {}
27
+ if tools is not None:
28
+ extra_kwargs["tools"] = tools
29
+ if reasoning_effort is not None:
30
+ extra_kwargs["reasoning_effort"] = reasoning_effort
31
+
32
+ if supports_chat_template(tokenizer):
33
+ # chat.py と同じfallbackチェーン
34
+ prompt = None
35
+ for kwargs in [extra_kwargs, {k: v for k, v in extra_kwargs.items() if k == "tools"}, {}]:
36
+ try:
37
+ prompt = tokenizer.apply_chat_template(
38
+ messages,
39
+ add_generation_prompt=False,
40
+ tokenize=False,
41
+ **kwargs,
42
+ )
43
+ break
44
+ except TypeError:
45
+ continue
46
+
47
+ if prompt is None:
48
+ prompt = str(messages)
49
+ else:
50
+ prompt = generate_merged_prompt(messages, capabilities)
51
+
52
+ # トークン化
53
+ add_special = tokenizer.bos_token is None or not prompt.startswith(
54
+ tokenizer.bos_token or ""
55
+ )
56
+ token_ids = tokenizer.encode(prompt, add_special_tokens=add_special)
57
+
58
+ result["token_ids"] = token_ids
59
+ result["token_count"] = len(token_ids)
60
+ except Exception as e:
61
+ result["error"] = str(e)
62
+
63
+ print(json.dumps(result), end="\0", flush=True)
@@ -0,0 +1,37 @@
1
+ [project]
2
+ name = "pytorch_driver"
3
+ version = "0.1.0"
4
+ description = "PyTorch (Transformers) driver for modular-prompt — CUDA runtime"
5
+ requires-python = ">=3.10,<3.14"
6
+ dependencies = [
7
+ "safetensors==0.8.0",
8
+ "tokenizers==0.22.2",
9
+ "transformers>=5.14.0",
10
+ ]
11
+
12
+ [dependency-groups]
13
+ dev = ["pytest>=9.0"]
14
+
15
+ [build-system]
16
+ requires = ["setuptools>=61.0"]
17
+ build-backend = "setuptools.build_meta"
18
+
19
+ [tool.pytest.ini_options]
20
+ testpaths = ["tests"]
21
+
22
+ [tool.setuptools]
23
+ py-modules = ["__main__", "server"]
24
+
25
+ [tool.setuptools.packages.find]
26
+ where = ["."]
27
+ include = ["backends*", "handlers*", "utils*"]
28
+
29
+ # setup-pytorch --variant cuda は CUDA wheel を明示インストールする。
30
+ # デフォルトは cu124。別の index を使う場合は --cuda <version> を指定する。
31
+ [[tool.uv.index]]
32
+ name = "pytorch-cuda"
33
+ url = "https://download.pytorch.org/whl/cu124"
34
+ explicit = true
35
+
36
+ [tool.uv.sources]
37
+ torch = { index = "pytorch-cuda" }
@@ -0,0 +1,158 @@
1
+ """JSON-RPC風サーバー: stdin/stdoutベースのリクエストディスパッチ"""
2
+ import json
3
+ import sys
4
+
5
+ from backends.base import ModelBackend
6
+ from handlers import handle_cache_prefill, handle_capabilities, handle_completion, handle_format_test, handle_generate, handle_render, handle_tokenize
7
+ from handlers.cancel import request_cancel, reset_cancel
8
+
9
+
10
+ MAX_READ_LINES = 10000
11
+
12
+
13
+ def read():
14
+ lines = []
15
+ while True:
16
+ line = sys.stdin.readline()
17
+ if not line:
18
+ return None
19
+ lines.append(line)
20
+ if len(lines) > MAX_READ_LINES:
21
+ sys.stderr.write(f"Error: read buffer exceeded {MAX_READ_LINES} lines, discarding\n")
22
+ lines.clear()
23
+ continue
24
+ try:
25
+ return json.loads(''.join(lines))
26
+ except json.JSONDecodeError:
27
+ continue
28
+
29
+
30
+ def _is_valid_generate_prompt(prompt) -> bool:
31
+ """prompt は非空文字列または非空のトークン ID リスト"""
32
+ if isinstance(prompt, str):
33
+ return bool(prompt)
34
+ if isinstance(prompt, list):
35
+ return len(prompt) > 0
36
+ return False
37
+
38
+
39
+ class Server:
40
+ def __init__(self, backend: ModelBackend, capabilities: dict):
41
+ self.backend = backend
42
+ self.capabilities = capabilities
43
+
44
+ def run(self):
45
+ while True:
46
+ req = read()
47
+ if req is None:
48
+ break
49
+ self._dispatch(req)
50
+
51
+ def _error_response(self, message: str) -> None:
52
+ sys.stderr.write(f"Error: {message}\n")
53
+ print(json.dumps({"error": message}), end='\0', flush=True)
54
+
55
+ def _dispatch(self, req: dict):
56
+ method = req.get('method')
57
+ if not method:
58
+ self._error_response("'method' field is required")
59
+ return
60
+
61
+ if method == 'cancel':
62
+ request_cancel()
63
+ return
64
+
65
+ reset_cancel()
66
+
67
+ try:
68
+ if method == 'capabilities':
69
+ handle_capabilities(self.capabilities)
70
+
71
+ elif method == 'format_test':
72
+ messages = req.get('messages')
73
+ if not messages:
74
+ self._error_response("'messages' field is required for format_test method")
75
+ return
76
+ handle_format_test(self.backend, self.capabilities, messages, req.get('options', {}), req.get('tools'))
77
+
78
+ elif method == 'tokenize':
79
+ messages = req.get('messages')
80
+ if messages is None:
81
+ self._error_response("'messages' field is required for tokenize method")
82
+ return
83
+ handle_tokenize(
84
+ self.backend, self.capabilities, messages,
85
+ tools=req.get('tools'),
86
+ reasoning_effort=req.get('reasoning_effort'),
87
+ )
88
+
89
+ elif method == 'cache_prefill':
90
+ cache_path = req.get('cache_path')
91
+ messages = req.get('messages')
92
+ if not cache_path or not messages:
93
+ self._error_response("'cache_path' and 'messages' fields are required for cache_prefill")
94
+ return
95
+ handle_cache_prefill(
96
+ self.backend,
97
+ self.capabilities,
98
+ cache_path,
99
+ messages,
100
+ base_cache_path=req.get('base_cache_path'),
101
+ trim_to_tokens=req.get('trim_to_tokens'),
102
+ prefix_offsets=req.get('prefix_offsets'),
103
+ prefix_hashes=req.get('prefix_hashes'),
104
+ tools=req.get('tools'),
105
+ reasoning_effort=req.get('reasoning_effort'),
106
+ images=req.get('images'),
107
+ max_image_size=req.get('maxImageSize', 768),
108
+ )
109
+
110
+ elif method == 'render':
111
+ messages = req.get('messages')
112
+ if not messages:
113
+ self._error_response("'messages' field is required for render method")
114
+ return
115
+ handle_render(
116
+ self.backend,
117
+ messages,
118
+ options=req.get('options', {}),
119
+ tools=req.get('tools'),
120
+ reasoning_effort=req.get('reasoning_effort'),
121
+ )
122
+
123
+ elif method == 'generate':
124
+ prompt = req.get('prompt')
125
+ if not _is_valid_generate_prompt(prompt):
126
+ self._error_response("'prompt' field is required for generate method")
127
+ return
128
+ images = req.get('images', [])
129
+ handle_generate(
130
+ self.backend,
131
+ prompt,
132
+ options=req.get('options', {}),
133
+ images=images if images else None,
134
+ max_image_size=req.get('maxImageSize', 768),
135
+ primer=req.get('primer'),
136
+ cache_path=req.get('cache_path'),
137
+ cache_trim_tokens=req.get('cache_trim_tokens'),
138
+ )
139
+
140
+ elif method == 'completion':
141
+ prompt = req.get('prompt')
142
+ if not prompt:
143
+ self._error_response("'prompt' field is required for completion method")
144
+ return
145
+ images = req.get('images', [])
146
+ handle_completion(
147
+ self.backend,
148
+ prompt,
149
+ options=req.get('options', {}),
150
+ images=images if images else None,
151
+ max_image_size=req.get('maxImageSize', 768),
152
+ )
153
+
154
+ else:
155
+ self._error_response(f"Unknown method '{method}'")
156
+
157
+ except Exception as e:
158
+ self._error_response(f"Error processing request: {e}")
File without changes
@@ -0,0 +1,207 @@
1
+ import pytest
2
+
3
+ pytest.importorskip("torch")
4
+
5
+ import json
6
+ from types import SimpleNamespace
7
+
8
+ from handlers.cache import handle_cache_prefill
9
+ from handlers.generate import handle_generate
10
+
11
+
12
+ class _Tokenizer:
13
+ bos_token = "<bos>"
14
+ chat_template = "template"
15
+
16
+ def __init__(self):
17
+ self.calls = []
18
+
19
+ def apply_chat_template(self, messages, **kwargs):
20
+ self.calls.append((messages, kwargs))
21
+ return "prefix"
22
+
23
+
24
+ class _Cache:
25
+ pass
26
+
27
+
28
+ class _Backend:
29
+ model_kind = "lm"
30
+
31
+ def __init__(self, cache=None, multi_chunk=False, cache_write_tokens=0):
32
+ self.tokenizer = _Tokenizer()
33
+ self.cache = cache
34
+ self.multi_chunk = multi_chunk
35
+ self.pending_cache_write_tokens = cache_write_tokens
36
+ self.calls = []
37
+
38
+ def get_tokenizer(self):
39
+ return self.tokenizer
40
+
41
+ def cache_prefill(self, cache_path, prompt, **kwargs):
42
+ self.calls.append(("prefill", cache_path, prompt, kwargs))
43
+ self.pending_cache_write_tokens = 2
44
+ return {
45
+ "cache_path": cache_path,
46
+ "token_count": 2,
47
+ "cache_write_tokens": 2,
48
+ }
49
+
50
+ def consume_cache_write_tokens(self, cache_path):
51
+ self.calls.append(("consume-write", cache_path))
52
+ result = self.pending_cache_write_tokens
53
+ self.pending_cache_write_tokens = 0
54
+ return result
55
+
56
+ def load_cache_from_file(self, cache_path, **kwargs):
57
+ self.calls.append(("load", cache_path, kwargs))
58
+ return self.cache
59
+
60
+ def get_cache_offset(self, prompt_cache):
61
+ return 2 if prompt_cache is self.cache else 0
62
+
63
+ def tokenize_prompt(self, prompt):
64
+ assert prompt == "prefix suffix"
65
+ return [1, 2, 3]
66
+
67
+ def stream_generate(self, prompt, options, images=None, prompt_cache=None):
68
+ self.calls.append(("generate", prompt, options, images, prompt_cache))
69
+ cache_read_tokens = (
70
+ 2 if self.cache is not None and prompt_cache is self.cache else None
71
+ )
72
+ if self.multi_chunk:
73
+ for index, text in enumerate(("a", "b", "c"), start=1):
74
+ yield SimpleNamespace(
75
+ text=text,
76
+ prompt_tokens=3 if index == 1 else None,
77
+ generation_tokens=index,
78
+ cache_read_tokens=cache_read_tokens if index == 1 else None,
79
+ )
80
+ return
81
+ yield SimpleNamespace(
82
+ text="ok",
83
+ prompt_tokens=3,
84
+ generation_tokens=1,
85
+ cache_read_tokens=cache_read_tokens,
86
+ )
87
+
88
+
89
+ def _json_response(output):
90
+ return json.loads(output.split("\0", 1)[0])
91
+
92
+
93
+ def test_cache_prefill_renders_without_generation_prompt(capsys):
94
+ backend = _Backend()
95
+
96
+ handle_cache_prefill(
97
+ backend,
98
+ {"special_tokens": {}},
99
+ "memory://prefix",
100
+ [{"role": "user", "content": "hello"}],
101
+ )
102
+
103
+ assert _json_response(capsys.readouterr().out) == {
104
+ "cache_path": "memory://prefix",
105
+ "token_count": 2,
106
+ "cache_write_tokens": 2,
107
+ }
108
+ assert backend.calls == [
109
+ ("prefill", "memory://prefix", "prefix", {
110
+ "base_cache_path": None,
111
+ "trim_to_tokens": None,
112
+ "prefix_offsets": None,
113
+ "prefix_hashes": None,
114
+ "images": None,
115
+ "max_image_size": 768,
116
+ })
117
+ ]
118
+ assert backend.tokenizer.calls[0][1]["add_generation_prompt"] is False
119
+
120
+
121
+ def test_generate_loads_cache_and_only_generates_suffix(capsys):
122
+ cache = _Cache()
123
+ backend = _Backend(cache, cache_write_tokens=2)
124
+
125
+ handle_generate(
126
+ backend,
127
+ "prefix suffix",
128
+ options={"max_tokens": 1},
129
+ cache_path="memory://prefix",
130
+ )
131
+
132
+ output = capsys.readouterr().out
133
+ assert output.startswith("ok")
134
+ meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
135
+ assert meta == {
136
+ "prompt_tokens": 3,
137
+ "generation_tokens": 1,
138
+ "cache_read_tokens": 2,
139
+ "cache_write_tokens": 2,
140
+ "cache_loaded": True,
141
+ }
142
+ generate_call = next(call for call in backend.calls if call[0] == "generate")
143
+ assert generate_call[1] == [3]
144
+ assert generate_call[4] is cache
145
+ load_call = next(call for call in backend.calls if call[0] == "load")
146
+ assert load_call[2]["prompt"] == "prefix suffix"
147
+ assert load_call[2]["images"] is None
148
+ assert load_call[2]["max_image_size"] == 768
149
+ assert ("consume-write", "memory://prefix") in backend.calls
150
+
151
+
152
+ def test_generate_preserves_usage_meta_across_multiple_chunks(capsys):
153
+ cache = _Cache()
154
+ backend = _Backend(cache, multi_chunk=True, cache_write_tokens=2)
155
+
156
+ handle_generate(
157
+ backend,
158
+ "prefix suffix",
159
+ options={"max_tokens": 3},
160
+ cache_path="memory://prefix",
161
+ )
162
+
163
+ output = capsys.readouterr().out
164
+ assert output.startswith("abc")
165
+ meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
166
+ assert meta == {
167
+ "prompt_tokens": 3,
168
+ "generation_tokens": 3,
169
+ "cache_read_tokens": 2,
170
+ "cache_write_tokens": 2,
171
+ "cache_loaded": True,
172
+ }
173
+
174
+
175
+ def test_generate_uses_cold_path_for_missing_cache(capsys):
176
+ backend = _Backend(cache=None)
177
+
178
+ handle_generate(
179
+ backend,
180
+ "prefix suffix",
181
+ cache_path="memory://missing",
182
+ )
183
+
184
+ output = capsys.readouterr().out
185
+ meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
186
+ assert meta["cache_loaded"] is False
187
+ assert "cache_read_tokens" not in meta
188
+ assert "cache_write_tokens" not in meta
189
+ generate_call = next(call for call in backend.calls if call[0] == "generate")
190
+ assert generate_call[1] == "prefix suffix"
191
+ assert generate_call[4] is None
192
+
193
+
194
+ def test_generate_removes_cached_prefix_from_token_ids(capsys):
195
+ cache = _Cache()
196
+ backend = _Backend(cache)
197
+
198
+ handle_generate(
199
+ backend,
200
+ [1, 2, 3],
201
+ cache_path="memory://prefix",
202
+ )
203
+
204
+ output = capsys.readouterr().out
205
+ assert '"cache_loaded": true' in output
206
+ generate_call = next(call for call in backend.calls if call[0] == "generate")
207
+ assert generate_call[1] == [3]
@@ -0,0 +1,14 @@
1
+ from utils.token_utils import get_capabilities
2
+
3
+
4
+ class _Tokenizer:
5
+ apply_chat_template = None
6
+ chat_template = None
7
+ special_tokens_map = {}
8
+ added_tokens_encoder = {}
9
+
10
+
11
+ def test_capabilities_advertise_cache_prefill():
12
+ capabilities = get_capabilities(_Tokenizer())
13
+
14
+ assert "cache_prefill" in capabilities["methods"]