Repository navigation
Expand file tree
/
Copy path_worker_causality.py
More file actions
171 lines (152 loc) · 6.14 KB
/
Copy path_worker_causality.py
File metadata and controls
171 lines (152 loc) · 6.14 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
# Copyright (c) Microsoft Corporation. All rights reserved.
"""Optional diagnostic decoding; consumers still validate enclosing self identity."""
import json
import logging
import re
import unicodedata
from collections.abc import Callable
from typing import Any, TypeVar
T = TypeVar("T")
_LOG = logging.getLogger(__name__)
_UUID = re.compile(r"[0-9a-fA-F]{8}-(?:[0-9a-fA-F]{4}-){3}[0-9a-fA-F]{12}")
def _uuid(value: Any) -> bool:
return isinstance(value, str) and _UUID.fullmatch(value) is not None
def _text(value: Any) -> bool:
return (
isinstance(value, str)
and bool(value)
and len(value.encode("utf-8")) <= 256
and all(unicodedata.category(char) not in ("Cc", "Cs") for char in value)
)
def _fields(value: dict[str, Any], allowed: set[str]) -> bool:
return value.keys() <= allowed
def _reference(value: Any, event_type: str) -> bool:
return (
isinstance(value, dict)
and _fields(value, {"sessionId", "eventId", "agentId", "eventType", "provenance"})
and _text(value.get("sessionId"))
and _uuid(value.get("eventId"))
and ("agentId" not in value or _text(value["agentId"]))
and value.get("eventType") == event_type
and value.get("provenance") in ("native", "ahp_coordinator")
)
def _source(value: Any) -> bool:
if not isinstance(value, dict) or not isinstance(value.get("input"), dict):
return False
source_input = value["input"]
if (
not _fields(
value,
{
"input",
"admissions",
"captureComplete",
"completion",
"notification",
"admittedDuring",
},
)
or not _fields(source_input, {"queueItemId", "agentId", "sender", "senderBridges"})
or type(value.get("captureComplete")) is not bool
or not _uuid(source_input.get("queueItemId"))
or not _text(source_input.get("agentId"))
or not isinstance(value.get("admissions"), list)
or len(value["admissions"]) > 32
):
return False
if "sender" in source_input and not _reference(source_input["sender"], "tool.execution_start"):
return False
if "senderBridges" in source_input:
edges = source_input["senderBridges"]
if "sender" not in source_input or not isinstance(edges, list) or len(edges) > 32:
return False
if any(
not isinstance(edge, dict)
or not _fields(edge, {"source", "reported"})
or not _reference(edge.get("source"), "tool.execution_start")
or not _reference(edge.get("reported"), "tool.execution_start")
for edge in edges
):
return False
for admission in value["admissions"]:
if (
not isinstance(admission, dict)
or not _fields(admission, {"kind", "messageId", "event", "ahpTurnId"})
or admission.get("kind") not in ("queued_input", "system_continuation")
or not _text(admission.get("messageId"))
or ("ahpTurnId" in admission and not _uuid(admission["ahpTurnId"]))
or (
"event" in admission
and (
not _reference(admission["event"], "user.message")
or admission["event"].get("agentId") != source_input["agentId"]
)
)
):
return False
for field, event_type in [
("completion", "subagent.completed"),
("admittedDuring", "assistant.turn_start"),
]:
if field in value and not _reference(value[field], event_type):
return False
if "notification" in value:
notification = value["notification"]
if (
not isinstance(notification, dict)
or not _fields(notification, {"deliveryId", "event", "mode"})
or not _uuid(notification.get("deliveryId"))
or notification.get("mode") not in ("queued", "immediate")
or (
"event" in notification
and not _reference(notification["event"], "system.notification")
)
):
return False
return True
def _supported(value: Any) -> bool:
return (
isinstance(value, dict)
and _fields(value, {"version", "observationProvenance", "sources", "captureComplete"})
and type(value.get("version")) is int
and value["version"] == 1
and value.get("observationProvenance") in ("native", "ahp_coordinator")
and type(value.get("captureComplete")) is bool
and isinstance(value.get("sources"), list)
and len(value["sources"]) <= 32
and all(
_source(source) and (not value["captureComplete"] or source["captureComplete"])
for source in value["sources"]
)
)
def _small(value: Any, remaining: list[int], depth: int = 0) -> bool:
remaining[0] -= 1
if remaining[0] < 0 or depth > 32:
return False
if isinstance(value, str):
remaining[0] -= len(value)
elif isinstance(value, dict):
for key, item in value.items():
if not _small(key, remaining, depth + 1) or not _small(item, remaining, depth + 1):
return False
elif isinstance(value, list):
for item in value:
if not _small(item, remaining, depth + 1):
return False
return remaining[0] >= 0
def optional_worker_causality(value: Any, load: Callable[[], T]) -> T | None:
if value is None:
return None
try:
if not _small(value, [4096]) or not _supported(value):
raise ValueError("invalid optional worker metadata")
size = 0
encoder = json.JSONEncoder(ensure_ascii=False, separators=(",", ":"))
for chunk in encoder.iterencode({"workerCausality": value}):
size += len(chunk.encode("utf-8"))
if size > 4096:
raise ValueError("optional worker metadata budget")
return load()
except (AssertionError, AttributeError, KeyError, TypeError, ValueError, UnicodeError):
_LOG.warning("Ignoring invalid, unsupported or oversized workerCausality metadata")
return None