token-runtime 0.1.0a2__py3-none-any.whl
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.
- token_runtime/__init__.py +1 -0
- token_runtime/adapters.py +181 -0
- token_runtime/agent_integrations.py +114 -0
- token_runtime/anthropic_adapter.py +301 -0
- token_runtime/anthropic_conformance.py +336 -0
- token_runtime/anthropic_preserved_thinking_cert.py +213 -0
- token_runtime/benchmark.py +106 -0
- token_runtime/capabilities.py +111 -0
- token_runtime/cli.py +284 -0
- token_runtime/codex_recertification.py +115 -0
- token_runtime/compat_cert.py +462 -0
- token_runtime/compatibility.py +81 -0
- token_runtime/config.py +72 -0
- token_runtime/conformance.py +36 -0
- token_runtime/contracts.py +70 -0
- token_runtime/engine.py +136 -0
- token_runtime/feature_flags.py +78 -0
- token_runtime/gateway.py +151 -0
- token_runtime/gemini_adapter.py +260 -0
- token_runtime/gemini_conformance.py +351 -0
- token_runtime/integrations.py +272 -0
- token_runtime/metrics.py +79 -0
- token_runtime/model.py +39 -0
- token_runtime/openai_certification.py +310 -0
- token_runtime/planner.py +96 -0
- token_runtime/reducers.py +179 -0
- token_runtime/store.py +36 -0
- token_runtime/strategies.py +50 -0
- token_runtime/terminal_ui.py +96 -0
- token_runtime-0.1.0a2.dist-info/METADATA +238 -0
- token_runtime-0.1.0a2.dist-info/RECORD +35 -0
- token_runtime-0.1.0a2.dist-info/WHEEL +5 -0
- token_runtime-0.1.0a2.dist-info/entry_points.txt +2 -0
- token_runtime-0.1.0a2.dist-info/licenses/LICENSE +202 -0
- token_runtime-0.1.0a2.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import TYPE_CHECKING, Any, Mapping, Protocol, runtime_checkable
|
|
4
|
+
|
|
5
|
+
from .model import RequestEnvelope
|
|
6
|
+
|
|
7
|
+
if TYPE_CHECKING:
|
|
8
|
+
from .capabilities import CapabilityKey, CapabilityProfile
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@runtime_checkable
|
|
12
|
+
class ClientAdapterContract(Protocol):
|
|
13
|
+
@property
|
|
14
|
+
def client_id(self) -> str: ...
|
|
15
|
+
|
|
16
|
+
@property
|
|
17
|
+
def protocol_ids(self) -> tuple[str, ...]: ...
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@runtime_checkable
|
|
21
|
+
class ProtocolAdapterContract(Protocol):
|
|
22
|
+
@property
|
|
23
|
+
def protocol_id(self) -> str: ...
|
|
24
|
+
|
|
25
|
+
def parse(self, payload: Mapping[str, Any]) -> RequestEnvelope: ...
|
|
26
|
+
|
|
27
|
+
def serialize(self, envelope: RequestEnvelope) -> dict[str, Any]: ...
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@runtime_checkable
|
|
31
|
+
class CapabilityProviderContract(Protocol):
|
|
32
|
+
@property
|
|
33
|
+
def provider_id(self) -> str: ...
|
|
34
|
+
|
|
35
|
+
def profile_for(self, key: "CapabilityKey") -> "CapabilityProfile | None": ...
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@runtime_checkable
|
|
39
|
+
class AgentIntegrationContract(Protocol):
|
|
40
|
+
@property
|
|
41
|
+
def agent_id(self) -> str: ...
|
|
42
|
+
|
|
43
|
+
@property
|
|
44
|
+
def capability_key(self) -> "CapabilityKey": ...
|
|
45
|
+
|
|
46
|
+
@property
|
|
47
|
+
def endpoint(self) -> str: ...
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
@runtime_checkable
|
|
51
|
+
class TokenizerAdapterContract(Protocol):
|
|
52
|
+
@property
|
|
53
|
+
def tokenizer_id(self) -> str: ...
|
|
54
|
+
|
|
55
|
+
def estimate(self, text: str) -> int: ...
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
@runtime_checkable
|
|
59
|
+
class StrategyContract(Protocol):
|
|
60
|
+
@property
|
|
61
|
+
def strategy_id(self) -> str: ...
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
@runtime_checkable
|
|
65
|
+
class BenchmarkAdapterContract(Protocol):
|
|
66
|
+
@property
|
|
67
|
+
def benchmark_id(self) -> str: ...
|
|
68
|
+
|
|
69
|
+
@property
|
|
70
|
+
def assertion_ids(self) -> tuple[str, ...]: ...
|
token_runtime/engine.py
ADDED
|
@@ -0,0 +1,136 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import replace
|
|
4
|
+
import re
|
|
5
|
+
from typing import Iterable, Protocol
|
|
6
|
+
|
|
7
|
+
from .model import OptimizationDecision, OptimizationResult, RequestEnvelope
|
|
8
|
+
from .planner import ContextPlanner
|
|
9
|
+
from .reducers import JsonToolOutputReducer, RepeatedLineReducer, RetrievedDuplicateReducer, ToolOutputReducer
|
|
10
|
+
from .store import RecoveryStore
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
_TOKEN_RE = re.compile(r"\w+|[^\w\s]", re.UNICODE)
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def estimate_tokens(envelope: RequestEnvelope) -> int:
|
|
17
|
+
units = sum(len(_TOKEN_RE.findall(block.text)) for block in envelope.blocks)
|
|
18
|
+
return int(round(units * 1.33))
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class ReducerLike(Protocol):
|
|
22
|
+
def reduce(self, block, store: RecoveryStore, *, protected: bool): ...
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class OptimizationEngine:
|
|
26
|
+
def __init__(
|
|
27
|
+
self,
|
|
28
|
+
*,
|
|
29
|
+
planner: ContextPlanner,
|
|
30
|
+
store: RecoveryStore,
|
|
31
|
+
reducers: Iterable[ReducerLike] | None = None,
|
|
32
|
+
):
|
|
33
|
+
self.planner = planner
|
|
34
|
+
self.store = store
|
|
35
|
+
self.reducers = tuple(reducers) if reducers is not None else (
|
|
36
|
+
JsonToolOutputReducer(),
|
|
37
|
+
RepeatedLineReducer(),
|
|
38
|
+
RetrievedDuplicateReducer(),
|
|
39
|
+
ToolOutputReducer(),
|
|
40
|
+
)
|
|
41
|
+
|
|
42
|
+
def optimize(self, envelope: RequestEnvelope) -> OptimizationResult:
|
|
43
|
+
before = estimate_tokens(envelope)
|
|
44
|
+
plan = self.planner.plan(envelope)
|
|
45
|
+
if plan.decision is OptimizationDecision.BYPASS:
|
|
46
|
+
return OptimizationResult(
|
|
47
|
+
decision=OptimizationDecision.BYPASS,
|
|
48
|
+
envelope=envelope,
|
|
49
|
+
reasons=plan.reasons,
|
|
50
|
+
before_estimated_tokens=before,
|
|
51
|
+
after_estimated_tokens=before,
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
blocks = []
|
|
55
|
+
changed_ids: list[str] = []
|
|
56
|
+
recovery_refs: list[str] = []
|
|
57
|
+
try:
|
|
58
|
+
for original in envelope.blocks:
|
|
59
|
+
current = original
|
|
60
|
+
for reducer in self.reducers:
|
|
61
|
+
reduction = reducer.reduce(
|
|
62
|
+
current,
|
|
63
|
+
self.store,
|
|
64
|
+
protected=original.id in plan.protected_ids,
|
|
65
|
+
)
|
|
66
|
+
current = reduction.block
|
|
67
|
+
if reduction.changed:
|
|
68
|
+
if original.id not in changed_ids:
|
|
69
|
+
changed_ids.append(original.id)
|
|
70
|
+
if reduction.recovery_ref and reduction.recovery_ref not in recovery_refs:
|
|
71
|
+
recovery_refs.append(reduction.recovery_ref)
|
|
72
|
+
blocks.append(current)
|
|
73
|
+
except Exception:
|
|
74
|
+
return self._fallback(envelope, before, "fallback_internal_error")
|
|
75
|
+
|
|
76
|
+
blocks, dedup_changed, dedup_refs = self._dedup_exact_evidence(
|
|
77
|
+
envelope, blocks, plan.protected_ids
|
|
78
|
+
)
|
|
79
|
+
for block_id in dedup_changed:
|
|
80
|
+
if block_id not in changed_ids:
|
|
81
|
+
changed_ids.append(block_id)
|
|
82
|
+
for ref in dedup_refs:
|
|
83
|
+
if ref not in recovery_refs:
|
|
84
|
+
recovery_refs.append(ref)
|
|
85
|
+
|
|
86
|
+
optimized = replace(envelope, blocks=tuple(blocks))
|
|
87
|
+
after = estimate_tokens(optimized)
|
|
88
|
+
if not changed_ids or after >= before:
|
|
89
|
+
return self._fallback(envelope, before, "no_eligible_reduction")
|
|
90
|
+
return OptimizationResult(
|
|
91
|
+
decision=OptimizationDecision.OPTIMIZE,
|
|
92
|
+
envelope=optimized,
|
|
93
|
+
reasons=("optimized",),
|
|
94
|
+
before_estimated_tokens=before,
|
|
95
|
+
after_estimated_tokens=after,
|
|
96
|
+
changed_block_ids=tuple(changed_ids),
|
|
97
|
+
recovery_refs=tuple(recovery_refs),
|
|
98
|
+
)
|
|
99
|
+
|
|
100
|
+
def _dedup_exact_evidence(self, envelope, blocks, protected_ids):
|
|
101
|
+
evidence_kinds = {"retrieved", "tool_output"}
|
|
102
|
+
canonical: dict[str, tuple[int, str]] = {}
|
|
103
|
+
for index, (original, current) in enumerate(zip(envelope.blocks, blocks)):
|
|
104
|
+
if original.kind in evidence_kinds:
|
|
105
|
+
canonical[current.text] = (index, original.id)
|
|
106
|
+
|
|
107
|
+
out = []
|
|
108
|
+
changed: list[str] = []
|
|
109
|
+
refs: list[str] = []
|
|
110
|
+
for index, (original, current) in enumerate(zip(envelope.blocks, blocks)):
|
|
111
|
+
if original.kind not in evidence_kinds or original.id in protected_ids:
|
|
112
|
+
out.append(current)
|
|
113
|
+
continue
|
|
114
|
+
canonical_index, canonical_id = canonical[current.text]
|
|
115
|
+
if index == canonical_index:
|
|
116
|
+
out.append(current)
|
|
117
|
+
continue
|
|
118
|
+
ref = self.store.put(original.text.encode())
|
|
119
|
+
marker = f"[TOKEN_DUPLICATE_OF:{canonical_id};TOKEN_REF:{ref}]"
|
|
120
|
+
if len(marker) >= len(current.text):
|
|
121
|
+
out.append(current)
|
|
122
|
+
continue
|
|
123
|
+
out.append(replace(current, text=marker))
|
|
124
|
+
changed.append(original.id)
|
|
125
|
+
refs.append(ref)
|
|
126
|
+
return out, tuple(changed), tuple(refs)
|
|
127
|
+
|
|
128
|
+
@staticmethod
|
|
129
|
+
def _fallback(envelope: RequestEnvelope, before: int, reason: str) -> OptimizationResult:
|
|
130
|
+
return OptimizationResult(
|
|
131
|
+
decision=OptimizationDecision.BYPASS,
|
|
132
|
+
envelope=envelope,
|
|
133
|
+
reasons=(reason,),
|
|
134
|
+
before_estimated_tokens=before,
|
|
135
|
+
after_estimated_tokens=before,
|
|
136
|
+
)
|
|
@@ -0,0 +1,78 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
from enum import Enum
|
|
5
|
+
|
|
6
|
+
from .compatibility import CompatibilityState
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class FeatureFlagState(str, Enum):
|
|
10
|
+
DISABLED = "disabled"
|
|
11
|
+
SHADOW = "shadow"
|
|
12
|
+
CANARY = "canary"
|
|
13
|
+
ENABLED = "enabled"
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class FeatureEffect(str, Enum):
|
|
17
|
+
NONE = "none"
|
|
18
|
+
OBSERVE = "observe"
|
|
19
|
+
EXECUTE = "execute"
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
@dataclass(frozen=True, slots=True)
|
|
23
|
+
class FeatureFlag:
|
|
24
|
+
name: str
|
|
25
|
+
state: FeatureFlagState
|
|
26
|
+
|
|
27
|
+
def __post_init__(self) -> None:
|
|
28
|
+
if not isinstance(self.state, FeatureFlagState):
|
|
29
|
+
raise TypeError("state must be FeatureFlagState")
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class DuplicateFeatureFlagError(ValueError):
|
|
33
|
+
pass
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def permitted_effect(
|
|
37
|
+
flag_state: FeatureFlagState,
|
|
38
|
+
compatibility_state: CompatibilityState,
|
|
39
|
+
) -> FeatureEffect:
|
|
40
|
+
if flag_state is FeatureFlagState.DISABLED:
|
|
41
|
+
return FeatureEffect.NONE
|
|
42
|
+
if compatibility_state in {
|
|
43
|
+
CompatibilityState.BLOCKED,
|
|
44
|
+
CompatibilityState.UNSUPPORTED,
|
|
45
|
+
}:
|
|
46
|
+
return FeatureEffect.NONE
|
|
47
|
+
if flag_state is FeatureFlagState.SHADOW:
|
|
48
|
+
return FeatureEffect.OBSERVE
|
|
49
|
+
if flag_state is FeatureFlagState.CANARY:
|
|
50
|
+
if compatibility_state in {
|
|
51
|
+
CompatibilityState.CANARY,
|
|
52
|
+
CompatibilityState.CERTIFIED,
|
|
53
|
+
}:
|
|
54
|
+
return FeatureEffect.EXECUTE
|
|
55
|
+
return FeatureEffect.OBSERVE
|
|
56
|
+
if compatibility_state is CompatibilityState.CERTIFIED:
|
|
57
|
+
return FeatureEffect.EXECUTE
|
|
58
|
+
return FeatureEffect.OBSERVE
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
class FeatureFlagRegistry:
|
|
62
|
+
def __init__(self) -> None:
|
|
63
|
+
self._flags: dict[str, FeatureFlag] = {}
|
|
64
|
+
|
|
65
|
+
def register(self, flag: FeatureFlag) -> FeatureFlag:
|
|
66
|
+
existing = self._flags.get(flag.name)
|
|
67
|
+
if existing is None:
|
|
68
|
+
self._flags[flag.name] = flag
|
|
69
|
+
return flag
|
|
70
|
+
if existing == flag:
|
|
71
|
+
return existing
|
|
72
|
+
raise DuplicateFeatureFlagError(f"conflicting feature flag: {flag.name}")
|
|
73
|
+
|
|
74
|
+
def get(self, name: str) -> FeatureFlag:
|
|
75
|
+
return self._flags.get(name, FeatureFlag(name, FeatureFlagState.DISABLED))
|
|
76
|
+
|
|
77
|
+
def snapshot(self) -> tuple[FeatureFlag, ...]:
|
|
78
|
+
return tuple(self._flags[name] for name in sorted(self._flags))
|
token_runtime/gateway.py
ADDED
|
@@ -0,0 +1,151 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
import json
|
|
5
|
+
import re
|
|
6
|
+
import time
|
|
7
|
+
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
8
|
+
from typing import Any
|
|
9
|
+
from urllib.error import HTTPError
|
|
10
|
+
from urllib.request import Request, urlopen
|
|
11
|
+
|
|
12
|
+
from .adapters import ChatCompletionsAdapter, ResponsesAdapter
|
|
13
|
+
from .metrics import MetricsStore
|
|
14
|
+
from .model import OptimizationDecision
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
_SAFE_LABEL_RE = re.compile(r"^[A-Za-z0-9._:/-]{1,128}$")
|
|
18
|
+
_HOP_HEADERS = {"host", "content-length", "connection", "transfer-encoding"}
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
@dataclass(frozen=True, slots=True)
|
|
22
|
+
class PreparedRequest:
|
|
23
|
+
body: bytes
|
|
24
|
+
optimized: bool
|
|
25
|
+
reasons: tuple[str, ...]
|
|
26
|
+
adapter: str | None = None
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class GatewayCore:
|
|
30
|
+
def __init__(self, *, engine, metrics: MetricsStore):
|
|
31
|
+
self.engine = engine
|
|
32
|
+
self.metrics = metrics
|
|
33
|
+
|
|
34
|
+
def prepare(self, path: str, raw_body: bytes) -> PreparedRequest:
|
|
35
|
+
started = time.perf_counter()
|
|
36
|
+
adapter = self._adapter(path)
|
|
37
|
+
if adapter is None:
|
|
38
|
+
prepared = PreparedRequest(raw_body, False, ("unsupported_endpoint",), None)
|
|
39
|
+
self._record(prepared, None, None, None, started)
|
|
40
|
+
return prepared
|
|
41
|
+
|
|
42
|
+
try:
|
|
43
|
+
payload = json.loads(raw_body.decode("utf-8"))
|
|
44
|
+
if not isinstance(payload, dict):
|
|
45
|
+
raise ValueError("JSON request must be an object")
|
|
46
|
+
except (UnicodeDecodeError, json.JSONDecodeError, ValueError):
|
|
47
|
+
prepared = PreparedRequest(raw_body, False, ("invalid_json",), adapter.name)
|
|
48
|
+
self._record(prepared, None, None, None, started)
|
|
49
|
+
return prepared
|
|
50
|
+
|
|
51
|
+
try:
|
|
52
|
+
envelope = adapter.parse(payload)
|
|
53
|
+
result = self.engine.optimize(envelope)
|
|
54
|
+
if result.decision is OptimizationDecision.OPTIMIZE:
|
|
55
|
+
serialized = adapter.serialize(result.envelope)
|
|
56
|
+
body = json.dumps(
|
|
57
|
+
serialized, ensure_ascii=False, separators=(",", ":")
|
|
58
|
+
).encode("utf-8")
|
|
59
|
+
prepared = PreparedRequest(body, True, result.reasons, adapter.name)
|
|
60
|
+
else:
|
|
61
|
+
prepared = PreparedRequest(raw_body, False, result.reasons, adapter.name)
|
|
62
|
+
self._record(prepared, result, payload, adapter.name, started)
|
|
63
|
+
return prepared
|
|
64
|
+
except Exception:
|
|
65
|
+
prepared = PreparedRequest(
|
|
66
|
+
raw_body, False, ("fallback_internal_error",), adapter.name
|
|
67
|
+
)
|
|
68
|
+
self._record(prepared, None, payload, adapter.name, started)
|
|
69
|
+
return prepared
|
|
70
|
+
|
|
71
|
+
@staticmethod
|
|
72
|
+
def _adapter(path: str):
|
|
73
|
+
clean_path = path.split("?", 1)[0]
|
|
74
|
+
if clean_path == "/v1/responses":
|
|
75
|
+
return ResponsesAdapter()
|
|
76
|
+
if clean_path == "/v1/chat/completions":
|
|
77
|
+
return ChatCompletionsAdapter()
|
|
78
|
+
return None
|
|
79
|
+
|
|
80
|
+
def _record(self, prepared, result, payload, adapter_name, started) -> None:
|
|
81
|
+
before = getattr(result, "before_estimated_tokens", None)
|
|
82
|
+
after = getattr(result, "after_estimated_tokens", None)
|
|
83
|
+
model = payload.get("model") if isinstance(payload, dict) else None
|
|
84
|
+
if not isinstance(model, str) or not _SAFE_LABEL_RE.fullmatch(model):
|
|
85
|
+
model = None
|
|
86
|
+
event = {
|
|
87
|
+
"decision": "optimize" if prepared.optimized else "bypass",
|
|
88
|
+
"adapter": adapter_name,
|
|
89
|
+
"model": model,
|
|
90
|
+
"before_input_tokens": before,
|
|
91
|
+
"after_input_tokens": after,
|
|
92
|
+
"latency_ms": round((time.perf_counter() - started) * 1000, 3),
|
|
93
|
+
"reason_codes": prepared.reasons,
|
|
94
|
+
"reducer_ids": (),
|
|
95
|
+
}
|
|
96
|
+
self.metrics.record(event)
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def _request_headers(handler: BaseHTTPRequestHandler) -> dict[str, str]:
|
|
100
|
+
return {
|
|
101
|
+
key: value
|
|
102
|
+
for key, value in handler.headers.items()
|
|
103
|
+
if key.lower() not in _HOP_HEADERS
|
|
104
|
+
}
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def _write_upstream_response(handler, response) -> None:
|
|
108
|
+
handler.send_response(response.status)
|
|
109
|
+
for key, value in response.headers.items():
|
|
110
|
+
if key.lower() not in _HOP_HEADERS:
|
|
111
|
+
handler.send_header(key, value)
|
|
112
|
+
handler.end_headers()
|
|
113
|
+
while True:
|
|
114
|
+
chunk = response.read(64 * 1024)
|
|
115
|
+
if not chunk:
|
|
116
|
+
break
|
|
117
|
+
handler.wfile.write(chunk)
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def build_server(host: str, port: int, *, upstream: str, core: GatewayCore):
|
|
121
|
+
upstream_base = upstream.rstrip("/")
|
|
122
|
+
|
|
123
|
+
class Handler(BaseHTTPRequestHandler):
|
|
124
|
+
def do_POST(self):
|
|
125
|
+
length = int(self.headers.get("Content-Length", "0"))
|
|
126
|
+
raw_body = self.rfile.read(length)
|
|
127
|
+
prepared = core.prepare(self.path, raw_body)
|
|
128
|
+
request = Request(
|
|
129
|
+
upstream_base + self.path,
|
|
130
|
+
data=prepared.body,
|
|
131
|
+
headers=_request_headers(self),
|
|
132
|
+
method="POST",
|
|
133
|
+
)
|
|
134
|
+
try:
|
|
135
|
+
with urlopen(request, timeout=60) as response:
|
|
136
|
+
_write_upstream_response(self, response)
|
|
137
|
+
except HTTPError as response:
|
|
138
|
+
_write_upstream_response(self, response)
|
|
139
|
+
|
|
140
|
+
def log_message(self, *args):
|
|
141
|
+
return
|
|
142
|
+
|
|
143
|
+
return ThreadingHTTPServer((host, port), Handler)
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def serve(host: str, port: int, *, upstream: str, core: GatewayCore) -> None:
|
|
147
|
+
server = build_server(host, port, upstream=upstream, core=core)
|
|
148
|
+
try:
|
|
149
|
+
server.serve_forever()
|
|
150
|
+
finally:
|
|
151
|
+
server.server_close()
|
|
@@ -0,0 +1,260 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from copy import deepcopy
|
|
4
|
+
import json
|
|
5
|
+
from typing import Any, Mapping, Sequence
|
|
6
|
+
|
|
7
|
+
from .model import ContextBlock, RequestEnvelope
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def _set_path(root: Any, path: Sequence[Any], value: str) -> None:
|
|
11
|
+
target = root
|
|
12
|
+
for key in path[:-1]:
|
|
13
|
+
target = target[key]
|
|
14
|
+
target[path[-1]] = value
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _block(
|
|
18
|
+
block_id: str,
|
|
19
|
+
kind: str,
|
|
20
|
+
text: str,
|
|
21
|
+
*,
|
|
22
|
+
role: str | None = None,
|
|
23
|
+
turn: int = 0,
|
|
24
|
+
path: Sequence[Any] | None = None,
|
|
25
|
+
) -> ContextBlock:
|
|
26
|
+
metadata = {} if path is None else {"path": tuple(path)}
|
|
27
|
+
return ContextBlock(
|
|
28
|
+
block_id,
|
|
29
|
+
kind,
|
|
30
|
+
text,
|
|
31
|
+
role=role,
|
|
32
|
+
turn_index=turn,
|
|
33
|
+
metadata=metadata,
|
|
34
|
+
)
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def _encoded(value: Mapping[str, Any]) -> str:
|
|
38
|
+
return json.dumps(dict(value), sort_keys=True, separators=(",", ":"))
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class GeminiGenerateContentAdapter:
|
|
42
|
+
protocol_id = "gemini_generate_content"
|
|
43
|
+
|
|
44
|
+
def serialize(self, envelope: RequestEnvelope) -> dict[str, Any]:
|
|
45
|
+
payload = deepcopy(envelope.opaque["original_payload"])
|
|
46
|
+
for block in envelope.blocks:
|
|
47
|
+
path = block.metadata.get("path")
|
|
48
|
+
if path:
|
|
49
|
+
_set_path(payload, path, block.text)
|
|
50
|
+
return payload
|
|
51
|
+
|
|
52
|
+
def parse(self, payload: Mapping[str, Any]) -> RequestEnvelope:
|
|
53
|
+
original = deepcopy(dict(payload))
|
|
54
|
+
blocks: list[ContextBlock] = []
|
|
55
|
+
safe = True
|
|
56
|
+
|
|
57
|
+
if "cachedContent" in payload:
|
|
58
|
+
cached_content = payload.get("cachedContent")
|
|
59
|
+
safe = (
|
|
60
|
+
isinstance(cached_content, str)
|
|
61
|
+
and cached_content.startswith("cachedContents/")
|
|
62
|
+
and len(cached_content) > len("cachedContents/")
|
|
63
|
+
)
|
|
64
|
+
|
|
65
|
+
system_blocks, system_safe = self._parse_system(payload.get("systemInstruction"))
|
|
66
|
+
blocks.extend(system_blocks)
|
|
67
|
+
safe = safe and system_safe
|
|
68
|
+
|
|
69
|
+
tool_blocks, tools_safe = self._parse_tools(payload.get("tools"))
|
|
70
|
+
blocks.extend(tool_blocks)
|
|
71
|
+
safe = safe and tools_safe
|
|
72
|
+
|
|
73
|
+
contents = payload.get("contents")
|
|
74
|
+
if not isinstance(contents, list) or not contents:
|
|
75
|
+
safe = False
|
|
76
|
+
contents = []
|
|
77
|
+
|
|
78
|
+
for index, content in enumerate(contents):
|
|
79
|
+
content_blocks, content_safe = self._parse_content(index, content)
|
|
80
|
+
blocks.extend(content_blocks)
|
|
81
|
+
safe = safe and content_safe
|
|
82
|
+
|
|
83
|
+
return RequestEnvelope(
|
|
84
|
+
blocks=tuple(blocks),
|
|
85
|
+
opaque={"original_payload": original},
|
|
86
|
+
adapter=self.protocol_id,
|
|
87
|
+
wire_safe=safe,
|
|
88
|
+
)
|
|
89
|
+
|
|
90
|
+
@staticmethod
|
|
91
|
+
def _parse_system(system: Any) -> tuple[list[ContextBlock], bool]:
|
|
92
|
+
if system is None:
|
|
93
|
+
return [], True
|
|
94
|
+
if not isinstance(system, Mapping):
|
|
95
|
+
return [], False
|
|
96
|
+
parts = system.get("parts")
|
|
97
|
+
if not isinstance(parts, list) or not parts:
|
|
98
|
+
return [], False
|
|
99
|
+
|
|
100
|
+
blocks: list[ContextBlock] = []
|
|
101
|
+
safe = True
|
|
102
|
+
for index, part in enumerate(parts):
|
|
103
|
+
if not isinstance(part, Mapping) or not isinstance(part.get("text"), str):
|
|
104
|
+
safe = False
|
|
105
|
+
continue
|
|
106
|
+
if set(part) != {"text"}:
|
|
107
|
+
safe = False
|
|
108
|
+
continue
|
|
109
|
+
blocks.append(
|
|
110
|
+
_block(
|
|
111
|
+
f"system-{index}",
|
|
112
|
+
"system",
|
|
113
|
+
part["text"],
|
|
114
|
+
role="system",
|
|
115
|
+
path=("systemInstruction", "parts", index, "text"),
|
|
116
|
+
)
|
|
117
|
+
)
|
|
118
|
+
return blocks, safe
|
|
119
|
+
|
|
120
|
+
@staticmethod
|
|
121
|
+
def _parse_tools(tools: Any) -> tuple[list[ContextBlock], bool]:
|
|
122
|
+
if tools is None:
|
|
123
|
+
return [], True
|
|
124
|
+
if not isinstance(tools, list):
|
|
125
|
+
return [], False
|
|
126
|
+
blocks: list[ContextBlock] = []
|
|
127
|
+
safe = True
|
|
128
|
+
for index, tool in enumerate(tools):
|
|
129
|
+
if not isinstance(tool, Mapping):
|
|
130
|
+
safe = False
|
|
131
|
+
continue
|
|
132
|
+
declarations = tool.get("functionDeclarations")
|
|
133
|
+
valid = (
|
|
134
|
+
set(tool) == {"functionDeclarations"}
|
|
135
|
+
and isinstance(declarations, list)
|
|
136
|
+
and bool(declarations)
|
|
137
|
+
and all(
|
|
138
|
+
isinstance(declaration, Mapping)
|
|
139
|
+
and isinstance(declaration.get("name"), str)
|
|
140
|
+
for declaration in declarations
|
|
141
|
+
)
|
|
142
|
+
)
|
|
143
|
+
blocks.append(_block(f"tool-schema-{index}", "tool_schema", _encoded(tool)))
|
|
144
|
+
safe = safe and valid
|
|
145
|
+
return blocks, safe
|
|
146
|
+
|
|
147
|
+
def _parse_content(
|
|
148
|
+
self,
|
|
149
|
+
index: int,
|
|
150
|
+
content: Any,
|
|
151
|
+
) -> tuple[list[ContextBlock], bool]:
|
|
152
|
+
if not isinstance(content, Mapping):
|
|
153
|
+
return [], False
|
|
154
|
+
role = content.get("role")
|
|
155
|
+
parts = content.get("parts")
|
|
156
|
+
if not isinstance(parts, list) or not parts:
|
|
157
|
+
return [], False
|
|
158
|
+
if role is None:
|
|
159
|
+
roleless_text = all(
|
|
160
|
+
isinstance(part, Mapping)
|
|
161
|
+
and set(part) == {"text"}
|
|
162
|
+
and isinstance(part.get("text"), str)
|
|
163
|
+
for part in parts
|
|
164
|
+
)
|
|
165
|
+
if not roleless_text:
|
|
166
|
+
return [], False
|
|
167
|
+
role = "user"
|
|
168
|
+
if role not in {"user", "model", "function"}:
|
|
169
|
+
return [], False
|
|
170
|
+
|
|
171
|
+
blocks: list[ContextBlock] = []
|
|
172
|
+
safe = True
|
|
173
|
+
for part_index, part in enumerate(parts):
|
|
174
|
+
part_blocks, part_safe = self._parse_part(index, part_index, role, part)
|
|
175
|
+
blocks.extend(part_blocks)
|
|
176
|
+
safe = safe and part_safe
|
|
177
|
+
return blocks, safe
|
|
178
|
+
|
|
179
|
+
@staticmethod
|
|
180
|
+
def _parse_part(
|
|
181
|
+
content_index: int,
|
|
182
|
+
part_index: int,
|
|
183
|
+
role: str,
|
|
184
|
+
part: Any,
|
|
185
|
+
) -> tuple[list[ContextBlock], bool]:
|
|
186
|
+
if not isinstance(part, Mapping):
|
|
187
|
+
return [], False
|
|
188
|
+
|
|
189
|
+
if "thoughtSignature" in part or "thought" in part:
|
|
190
|
+
return [
|
|
191
|
+
_block(
|
|
192
|
+
f"content-{content_index}-protocol-{part_index}",
|
|
193
|
+
"protocol_state",
|
|
194
|
+
_encoded(part),
|
|
195
|
+
role="assistant" if role == "model" else role,
|
|
196
|
+
turn=content_index,
|
|
197
|
+
)
|
|
198
|
+
], False
|
|
199
|
+
|
|
200
|
+
if "text" in part:
|
|
201
|
+
text = part.get("text")
|
|
202
|
+
if role not in {"user", "model"} or not isinstance(text, str):
|
|
203
|
+
return [], False
|
|
204
|
+
if set(part) != {"text"}:
|
|
205
|
+
return [], False
|
|
206
|
+
kind = "assistant" if role == "model" else "user"
|
|
207
|
+
return [
|
|
208
|
+
_block(
|
|
209
|
+
f"content-{content_index}-part-{part_index}",
|
|
210
|
+
kind,
|
|
211
|
+
text,
|
|
212
|
+
role=kind,
|
|
213
|
+
turn=content_index,
|
|
214
|
+
path=("contents", content_index, "parts", part_index, "text"),
|
|
215
|
+
)
|
|
216
|
+
], True
|
|
217
|
+
|
|
218
|
+
if "functionCall" in part:
|
|
219
|
+
call = part.get("functionCall")
|
|
220
|
+
valid = (
|
|
221
|
+
set(part) == {"functionCall"}
|
|
222
|
+
and role == "model"
|
|
223
|
+
and isinstance(call, Mapping)
|
|
224
|
+
and isinstance(call.get("name"), str)
|
|
225
|
+
and ("id" not in call or isinstance(call.get("id"), str))
|
|
226
|
+
and ("args" not in call or isinstance(call.get("args"), Mapping))
|
|
227
|
+
)
|
|
228
|
+
return [
|
|
229
|
+
_block(
|
|
230
|
+
f"content-{content_index}-function-call-{part_index}",
|
|
231
|
+
"tool_call",
|
|
232
|
+
_encoded(part),
|
|
233
|
+
role="assistant",
|
|
234
|
+
turn=content_index,
|
|
235
|
+
)
|
|
236
|
+
], valid
|
|
237
|
+
|
|
238
|
+
if "functionResponse" in part:
|
|
239
|
+
response = part.get("functionResponse")
|
|
240
|
+
valid = (
|
|
241
|
+
set(part) == {"functionResponse"}
|
|
242
|
+
and role in {"user", "function"}
|
|
243
|
+
and isinstance(response, Mapping)
|
|
244
|
+
and isinstance(response.get("name"), str)
|
|
245
|
+
and isinstance(response.get("response"), Mapping)
|
|
246
|
+
and ("id" not in response or isinstance(response.get("id"), str))
|
|
247
|
+
and not response.get("parts")
|
|
248
|
+
)
|
|
249
|
+
kind = "tool_output" if valid else "protocol_state"
|
|
250
|
+
return [
|
|
251
|
+
_block(
|
|
252
|
+
f"content-{content_index}-function-response-{part_index}",
|
|
253
|
+
kind,
|
|
254
|
+
_encoded(part),
|
|
255
|
+
role="tool",
|
|
256
|
+
turn=content_index,
|
|
257
|
+
)
|
|
258
|
+
], False
|
|
259
|
+
|
|
260
|
+
return [], False
|