autoinference-utils 0.2.5__tar.gz → 0.2.6__tar.gz

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.
@@ -12,6 +12,7 @@ node_modules/
12
12
  build/
13
13
  develop-eggs/
14
14
  dist/
15
+ !.github/actions/paths-filter/dist/
15
16
  downloads/
16
17
  eggs/
17
18
  .eggs/
@@ -151,6 +152,7 @@ activemq-data/
151
152
  .env
152
153
  .envrc
153
154
  .venv
155
+ .autoinference/
154
156
  env/
155
157
  venv/
156
158
  ENV/
@@ -239,10 +241,6 @@ Network Trash Folder
239
241
  Temporary Items
240
242
  .apdisk
241
243
 
244
+ # Ephemeral files from Experiments
242
245
  benchmark_results/
243
- docs/*
244
- !docs/data-model.md
245
- !docs/dogfood-feedback.md
246
- !docs/main-architecture-explainer.md
247
- !docs/recipe-seeding.md
248
- !docs/sweep-learnings.md
246
+ src/autoinference/deployments/
@@ -1,6 +1,6 @@
1
- Metadata-Version: 2.4
1
+ Metadata-Version: 2.5
2
2
  Name: autoinference-utils
3
- Version: 0.2.5
3
+ Version: 0.2.6
4
4
  Summary: Shared endpoint abstractions for autoinference deployments
5
5
  Requires-Python: >=3.10
6
6
  Description-Content-Type: text/markdown
@@ -30,12 +30,55 @@ BENCH_MODE_PORT_OFFSET = 10000
30
30
  MODAL_FLASH_REQUEST_UUID_HEADER= "modal-flash-request-uuid"
31
31
  MODAL_SESSION_ID_HEADER = "modal-session-id"
32
32
 
33
- # Endpoint class names whose servers emit OAI streaming chunks that
34
- # sglang.bench_serving can't parse. deploy_and_bench auto-injects
35
- # --openai-stream-compat for deployments importing any of these.
36
33
  ENDPOINTS_REQUIRING_OAI_STREAM_COMPAT: tuple[str, ...] = ("TRTLLMEndpoint",)
37
34
 
38
35
 
36
+ def server_arg_tokens(flag: str, value: str) -> list[str]:
37
+ """Render one server arg as argv tokens.
38
+
39
+ A value splits on whitespace, so a repeatable flag takes its values as one
40
+ entry. A lone value starting with `-` keeps the attached `--flag=value`
41
+ spelling, the only one an argparse-style parser does not read as a flag.
42
+ """
43
+ if value and "=" in flag:
44
+ raise ValueError(f"server arg {flag!r} already carries a value: {value!r}")
45
+ if value == "":
46
+ return [flag]
47
+ tokens = value.split()
48
+ if len(tokens) == 1 and tokens[0].startswith("-"):
49
+ return [f"{flag}={tokens[0]}"]
50
+ return [flag, *tokens]
51
+
52
+
53
+ def _flag_name(key: str) -> str:
54
+ """Return the flag a key sets, whose value may be attached as `--flag=value`."""
55
+ return key.partition("=")[0]
56
+
57
+
58
+ def get_server_arg(
59
+ server_args: Mapping[str, str], flag: str
60
+ ) -> tuple[str, str] | None:
61
+ """Return the key and value setting a flag under either spelling."""
62
+ for key, value in server_args.items():
63
+ if _flag_name(key) == flag:
64
+ return key, value
65
+ return None
66
+
67
+
68
+ def has_server_arg(server_args: Mapping[str, str], flag: str) -> bool:
69
+ """Report whether a flag is set under either spelling, plain or attached."""
70
+ return get_server_arg(server_args, flag) is not None
71
+
72
+
73
+ def _merge_server_args(
74
+ defaults: Mapping[str, str], overrides: Mapping[str, str]
75
+ ) -> dict[str, str]:
76
+ """Overlay overrides on defaults by flag name, so an attached spelling wins."""
77
+ merged = {_flag_name(key): (key, value) for key, value in defaults.items()}
78
+ merged.update((_flag_name(key), (key, value)) for key, value in overrides.items())
79
+ return dict(merged.values())
80
+
81
+
39
82
  class Endpoint(ABC):
40
83
  """A thing with a URL that you can start and stop."""
41
84
 
@@ -147,21 +190,27 @@ class SGLangEndpoint(Endpoint):
147
190
  if os.environ.get(DUMMY_WEIGHTS_ENV) == "1":
148
191
  if (
149
192
  self.load_format is None
150
- and "--load-format" not in self.extra_server_args
193
+ and not has_server_arg(self.extra_server_args, "--load-format")
151
194
  ):
152
195
  self.load_format = "dummy"
153
- self.extra_server_args.setdefault("--model-loader-extra-config", "{}")
154
-
196
+ if not has_server_arg(self.extra_server_args, "--model-loader-extra-config"):
197
+ self.extra_server_args["--model-loader-extra-config"] = "{}"
198
+
155
199
  # Request logging enabled, log only metadata by default
156
200
  if self.log_requests_level >= 0:
157
- # Enable request logging
158
- if "--log-requests" not in self.extra_server_args:
159
- self.extra_server_args.setdefault("--log-requests", "")
201
+ # Enable request logging
202
+ if not has_server_arg(self.extra_server_args, "--log-requests"):
203
+ self.extra_server_args["--log-requests"] = ""
160
204
  # Set request logging level
161
- if "--log-requests-level" not in self.extra_server_args:
162
- self.extra_server_args.setdefault("--log-requests-level", str(self.log_requests_level))
205
+ level = get_server_arg(self.extra_server_args, "--log-requests-level")
206
+ if level is None:
207
+ self.extra_server_args["--log-requests-level"] = str(self.log_requests_level)
163
208
  else:
164
- print(f"[endpoint] --log-requests-level already set to {self.extra_server_args['--log-requests-level']} in server args")
209
+ key, value = level
210
+ print(
211
+ "[endpoint] --log-requests-level already set in server args: "
212
+ f"{value or key}"
213
+ )
165
214
 
166
215
  # Add Modal-forwarded request headers to SGLang logged headers
167
216
  sglang_log_request_headers = [
@@ -221,12 +270,9 @@ class SGLangEndpoint(Endpoint):
221
270
  ]
222
271
  )
223
272
 
224
- merged = {**self.DEFAULT_OPERATIONAL_ARGS, **self.extra_server_args}
273
+ merged = _merge_server_args(self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args)
225
274
  for key, value in merged.items():
226
- if value == "":
227
- cmd.append(key)
228
- else:
229
- cmd.extend([key, *value.split()])
275
+ cmd.extend(server_arg_tokens(key, value))
230
276
 
231
277
  return cmd
232
278
 
@@ -297,7 +343,8 @@ class VLLMEndpoint(Endpoint):
297
343
  self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
298
344
 
299
345
  if os.environ.get(DUMMY_WEIGHTS_ENV) == "1":
300
- self.extra_server_args.setdefault("--load-format", "dummy")
346
+ if not has_server_arg(self.extra_server_args, "--load-format"):
347
+ self.extra_server_args["--load-format"] = "dummy"
301
348
 
302
349
  def _build_cmd(self) -> list[str]:
303
350
  cmd = [
@@ -307,10 +354,7 @@ class VLLMEndpoint(Endpoint):
307
354
  "--model", self.model,
308
355
  ]
309
356
  for key, value in self.extra_server_args.items():
310
- if value == "":
311
- cmd.append(key)
312
- else:
313
- cmd.extend([key, *value.split()])
357
+ cmd.extend(server_arg_tokens(key, value))
314
358
  return cmd
315
359
 
316
360
  def start(self):
@@ -377,10 +421,7 @@ class TRTLLMEndpoint(Endpoint):
377
421
  "--port", str(self.worker_port),
378
422
  ]
379
423
  for key, value in self.extra_server_args.items():
380
- if value == "":
381
- cmd.append(key)
382
- else:
383
- cmd.extend([key, *value.split()])
424
+ cmd.extend(server_arg_tokens(key, value))
384
425
  return cmd
385
426
 
386
427
  def health_check(self) -> str | None:
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "autoinference-utils"
3
- version = "0.2.5"
3
+ version = "0.2.6"
4
4
  description = "Shared endpoint abstractions for autoinference deployments"
5
5
  readme = "README.md"
6
6
  requires-python = ">=3.10"