autoinference-utils 0.2.3__tar.gz → 0.2.5__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.3
3
+ Version: 0.2.5
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,7 +27,7 @@ BENCH_MODE_ENV = "ENDPOINT_BENCH_MODE"
26
27
  DUMMY_WEIGHTS_ENV = "ENDPOINT_DUMMY"
27
28
  BENCH_MODE_PORT_OFFSET = 10000
28
29
 
29
- MODAL_FLASH_REQUEST_UUID_HEADER= "x-modal-flash-request-uuid"
30
+ MODAL_FLASH_REQUEST_UUID_HEADER= "modal-flash-request-uuid"
30
31
  MODAL_SESSION_ID_HEADER = "modal-session-id"
31
32
 
32
33
  # Endpoint class names whose servers emit OAI streaming chunks that
@@ -760,32 +761,86 @@ def warmup_chat_completions(
760
761
  request_headers.update(headers)
761
762
 
762
763
  for request_idx in range(successful_requests):
763
- for attempt in range(max_attempts_per_request):
764
- try:
765
- _post_json(
766
- url,
767
- payload=payload,
768
- headers=request_headers,
769
- timeout=request_timeout,
770
- )
771
- break
772
- except (
773
- urllib.error.HTTPError,
774
- urllib.error.URLError,
775
- TimeoutError,
776
- OSError,
777
- ) as exc:
778
- if attempt + 1 == max_attempts_per_request:
779
- detail = (
780
- _format_http_error(exc)
781
- if isinstance(exc, urllib.error.HTTPError)
782
- else f"{type(exc).__name__}: {exc}"
783
- )
784
- raise RuntimeError(
785
- f"warmup request {request_idx + 1}/"
786
- f"{successful_requests}: {detail}"
787
- ) from exc
788
- 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")
789
844
 
790
845
 
791
846
  def start_heartbeat_thread(
@@ -922,13 +977,51 @@ def _post_json(
922
977
  payload: Mapping[str, Any],
923
978
  headers: Mapping[str, str],
924
979
  timeout: float | None = None,
925
- ) -> int:
980
+ ) -> bytes:
926
981
  body = json.dumps(payload).encode("utf-8")
927
982
  req = urllib.request.Request(
928
983
  url, data=body, headers=dict(headers), method="POST"
929
984
  )
930
985
  with urllib.request.urlopen(req, timeout=timeout) as resp:
931
- 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")
932
1025
 
933
1026
 
934
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.3"
3
+ version = "0.2.5"
4
4
  description = "Shared endpoint abstractions for autoinference deployments"
5
5
  readme = "README.md"
6
6
  requires-python = ">=3.10"