autoinference-utils 0.2.2__tar.gz → 0.2.4__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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: autoinference-utils
3
- Version: 0.2.2
3
+ Version: 0.2.4
4
4
  Summary: Shared endpoint abstractions for autoinference deployments
5
5
  Requires-Python: >=3.10
6
6
  Description-Content-Type: text/markdown
@@ -11,6 +11,7 @@ from __future__ import annotations
11
11
 
12
12
  import http.server
13
13
  import json
14
+ import math
14
15
  import os
15
16
  import shlex
16
17
  import subprocess
@@ -26,6 +27,9 @@ BENCH_MODE_ENV = "ENDPOINT_BENCH_MODE"
26
27
  DUMMY_WEIGHTS_ENV = "ENDPOINT_DUMMY"
27
28
  BENCH_MODE_PORT_OFFSET = 10000
28
29
 
30
+ MODAL_FLASH_REQUEST_UUID_HEADER= "x-modal-flash-request-uuid"
31
+ MODAL_SESSION_ID_HEADER = "modal-session-id"
32
+
29
33
  # Endpoint class names whose servers emit OAI streaming chunks that
30
34
  # sglang.bench_serving can't parse. deploy_and_bench auto-injects
31
35
  # --openai-stream-compat for deployments importing any of these.
@@ -93,6 +97,7 @@ class SGLangEndpoint(Endpoint):
93
97
  health_timeout: float = 20 * 60,
94
98
  health_poll_interval: float = 5.0,
95
99
  health_request_timeout: float = 5.0,
100
+ log_requests_level: int = 0,
96
101
  ):
97
102
  # In bench mode SGLang runs on worker_port + 10000 and a /bench proxy
98
103
  # listens on the original worker_port; both fall through to the same
@@ -124,6 +129,13 @@ class SGLangEndpoint(Endpoint):
124
129
  self.health_timeout = health_timeout
125
130
  self.health_poll_interval = health_poll_interval
126
131
  self.health_request_timeout = health_request_timeout
132
+
133
+ # SGLang Log request level: -1 = disabled, 0 = metadata only, 1,2,3 in increasing verbosity
134
+ # Refer to https://docs.sglang.io/docs/advanced_features/server_arguments#logging for more details
135
+ if not (log_requests_level >= -1 and log_requests_level <= 3):
136
+ raise ValueError("log_requests_level must be between -1 and 3")
137
+ self.log_requests_level = log_requests_level
138
+
127
139
  self._proc: Optional[subprocess.Popen] = None
128
140
  self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
129
141
 
@@ -139,6 +151,32 @@ class SGLangEndpoint(Endpoint):
139
151
  ):
140
152
  self.load_format = "dummy"
141
153
  self.extra_server_args.setdefault("--model-loader-extra-config", "{}")
154
+
155
+ # Request logging enabled, log only metadata by default
156
+ 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", "")
160
+ # 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))
163
+ else:
164
+ print(f"[endpoint] --log-requests-level already set to {self.extra_server_args['--log-requests-level']} in server args")
165
+
166
+ # Add Modal-forwarded request headers to SGLang logged headers
167
+ sglang_log_request_headers = [
168
+ h.strip().lower()
169
+ for h in os.environ.get("SGLANG_LOG_REQUEST_HEADERS", "").split(",")
170
+ if h.strip()
171
+ ]
172
+ if MODAL_FLASH_REQUEST_UUID_HEADER not in sglang_log_request_headers:
173
+ sglang_log_request_headers.append(MODAL_FLASH_REQUEST_UUID_HEADER)
174
+ if MODAL_SESSION_ID_HEADER not in sglang_log_request_headers:
175
+ sglang_log_request_headers.append(MODAL_SESSION_ID_HEADER)
176
+
177
+ self.log_request_headers = ",".join(sglang_log_request_headers)
178
+ os.environ["SGLANG_LOG_REQUEST_HEADERS"] = self.log_request_headers
179
+
142
180
 
143
181
  def _build_cmd(self) -> list[str]:
144
182
  cmd = [
@@ -723,32 +761,86 @@ def warmup_chat_completions(
723
761
  request_headers.update(headers)
724
762
 
725
763
  for request_idx in range(successful_requests):
726
- for attempt in range(max_attempts_per_request):
727
- try:
728
- _post_json(
729
- url,
730
- payload=payload,
731
- headers=request_headers,
732
- timeout=request_timeout,
733
- )
734
- break
735
- except (
736
- urllib.error.HTTPError,
737
- urllib.error.URLError,
738
- TimeoutError,
739
- OSError,
740
- ) as exc:
741
- if attempt + 1 == max_attempts_per_request:
742
- detail = (
743
- _format_http_error(exc)
744
- if isinstance(exc, urllib.error.HTTPError)
745
- else f"{type(exc).__name__}: {exc}"
746
- )
747
- raise RuntimeError(
748
- f"warmup request {request_idx + 1}/"
749
- f"{successful_requests}: {detail}"
750
- ) from exc
751
- time.sleep(retry_delay)
764
+ _post_json_with_retries(
765
+ url,
766
+ payload=payload,
767
+ headers=request_headers,
768
+ request_timeout=request_timeout,
769
+ max_attempts=max_attempts_per_request,
770
+ retry_delay=retry_delay,
771
+ description=(
772
+ f"warmup request {request_idx + 1}/{successful_requests}"
773
+ ),
774
+ )
775
+
776
+
777
+ def validate_embeddings_endpoint(
778
+ *,
779
+ port: int,
780
+ payload: Mapping[str, Any],
781
+ headers: Mapping[str, str] | None = None,
782
+ request_timeout: float = 30.0,
783
+ max_attempts: int = 2,
784
+ retry_delay: float = 1.0,
785
+ ) -> None:
786
+ """Validate one OpenAI-compatible embeddings request."""
787
+ inputs = payload.get("input")
788
+ if isinstance(inputs, str):
789
+ expected_count = 1
790
+ elif isinstance(inputs, list) and inputs:
791
+ expected_count = (
792
+ 1 if all(isinstance(item, int) for item in inputs) else len(inputs)
793
+ )
794
+ else:
795
+ raise ValueError(
796
+ "embedding validation payload must contain non-empty input"
797
+ )
798
+
799
+ request_headers = {"Content-Type": "application/json"}
800
+ if headers:
801
+ request_headers.update(headers)
802
+ body = _post_json_with_retries(
803
+ f"http://127.0.0.1:{port}/v1/embeddings",
804
+ payload=payload,
805
+ headers=request_headers,
806
+ request_timeout=request_timeout,
807
+ max_attempts=max_attempts,
808
+ retry_delay=retry_delay,
809
+ description="embedding validation request",
810
+ )
811
+ try:
812
+ response = json.loads(body)
813
+ data = response["data"]
814
+ except (ValueError, KeyError, TypeError) as exc:
815
+ raise RuntimeError("embedding validation returned an invalid response") from exc
816
+
817
+ if not isinstance(data, list) or len(data) != expected_count:
818
+ raise RuntimeError(
819
+ f"embedding validation returned {len(data) if isinstance(data, list) else 0} "
820
+ f"embeddings; expected {expected_count}"
821
+ )
822
+
823
+ dimension: int | None = None
824
+ for index, item in enumerate(data):
825
+ embedding = item.get("embedding") if isinstance(item, Mapping) else None
826
+ if not isinstance(embedding, list) or not embedding:
827
+ raise RuntimeError(
828
+ f"embedding validation result {index} has no embedding vector"
829
+ )
830
+ if any(
831
+ isinstance(value, bool)
832
+ or not isinstance(value, (int, float))
833
+ or not math.isfinite(value)
834
+ for value in embedding
835
+ ):
836
+ raise RuntimeError(
837
+ f"embedding validation result {index} contains "
838
+ "non-numeric or non-finite values"
839
+ )
840
+ if dimension is None:
841
+ dimension = len(embedding)
842
+ elif len(embedding) != dimension:
843
+ raise RuntimeError("embedding validation returned inconsistent dimensions")
752
844
 
753
845
 
754
846
  def start_heartbeat_thread(
@@ -885,13 +977,51 @@ def _post_json(
885
977
  payload: Mapping[str, Any],
886
978
  headers: Mapping[str, str],
887
979
  timeout: float | None = None,
888
- ) -> int:
980
+ ) -> bytes:
889
981
  body = json.dumps(payload).encode("utf-8")
890
982
  req = urllib.request.Request(
891
983
  url, data=body, headers=dict(headers), method="POST"
892
984
  )
893
985
  with urllib.request.urlopen(req, timeout=timeout) as resp:
894
- return resp.getcode()
986
+ return resp.read()
987
+
988
+
989
+ def _post_json_with_retries(
990
+ url: str,
991
+ *,
992
+ payload: Mapping[str, Any],
993
+ headers: Mapping[str, str],
994
+ request_timeout: float,
995
+ max_attempts: int,
996
+ retry_delay: float,
997
+ description: str,
998
+ ) -> bytes:
999
+ if max_attempts < 1:
1000
+ raise ValueError("max_attempts must be at least 1")
1001
+
1002
+ for attempt in range(max_attempts):
1003
+ try:
1004
+ return _post_json(
1005
+ url,
1006
+ payload=payload,
1007
+ headers=headers,
1008
+ timeout=request_timeout,
1009
+ )
1010
+ except (
1011
+ urllib.error.HTTPError,
1012
+ urllib.error.URLError,
1013
+ TimeoutError,
1014
+ OSError,
1015
+ ) as exc:
1016
+ if attempt + 1 == max_attempts:
1017
+ detail = (
1018
+ _format_http_error(exc)
1019
+ if isinstance(exc, urllib.error.HTTPError)
1020
+ else f"{type(exc).__name__}: {exc}"
1021
+ )
1022
+ raise RuntimeError(f"{description}: {detail}") from exc
1023
+ time.sleep(retry_delay)
1024
+ raise AssertionError("unreachable")
895
1025
 
896
1026
 
897
1027
  def _format_http_error(exc: urllib.error.HTTPError) -> str:
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "autoinference-utils"
3
- version = "0.2.2"
3
+ version = "0.2.4"
4
4
  description = "Shared endpoint abstractions for autoinference deployments"
5
5
  readme = "README.md"
6
6
  requires-python = ">=3.10"