turnsafe 0.1.0__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.
- turnsafe/__init__.py +66 -0
- turnsafe/_match.py +126 -0
- turnsafe/_providers.py +351 -0
- turnsafe/_thinking.py +99 -0
- turnsafe/_types.py +119 -0
- turnsafe/fit.py +169 -0
- turnsafe/py.typed +0 -0
- turnsafe/repair.py +381 -0
- turnsafe/tokens.py +51 -0
- turnsafe/validate.py +196 -0
- turnsafe-0.1.0.dist-info/METADATA +361 -0
- turnsafe-0.1.0.dist-info/RECORD +14 -0
- turnsafe-0.1.0.dist-info/WHEEL +4 -0
- turnsafe-0.1.0.dist-info/licenses/LICENSE +21 -0
turnsafe/__init__.py
ADDED
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
"""turnsafe: validate, repair and trim LLM chat history without orphaning tool calls.
|
|
2
|
+
|
|
3
|
+
import turnsafe
|
|
4
|
+
|
|
5
|
+
report = turnsafe.validate(messages) # what the API would reject, and why
|
|
6
|
+
messages = turnsafe.repair(messages) # fix orphaned / missing / misplaced results
|
|
7
|
+
messages = turnsafe.fit(messages, 100_000) # trim whole turns to a token budget
|
|
8
|
+
|
|
9
|
+
Works on plain OpenAI Chat Completions and Anthropic Messages dicts. No
|
|
10
|
+
dependencies, no framework, no network.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from ._providers import detect as detect_provider
|
|
14
|
+
from ._types import (
|
|
15
|
+
DUPLICATE_TOOL_CALL_ID,
|
|
16
|
+
DUPLICATE_TOOL_RESULT,
|
|
17
|
+
EMPTY_CONTENT,
|
|
18
|
+
FIRST_MESSAGE_NOT_USER,
|
|
19
|
+
MISPLACED_TOOL_RESULT,
|
|
20
|
+
ORPHAN_TOOL_RESULT,
|
|
21
|
+
SYSTEM_IN_MESSAGES,
|
|
22
|
+
TOOL_RESULT_NOT_FIRST,
|
|
23
|
+
UNANSWERED_TOOL_CALL,
|
|
24
|
+
ContextBudgetError,
|
|
25
|
+
Fix,
|
|
26
|
+
InvalidHistoryError,
|
|
27
|
+
ProviderDetectionError,
|
|
28
|
+
Report,
|
|
29
|
+
TurnsafeError,
|
|
30
|
+
Violation,
|
|
31
|
+
)
|
|
32
|
+
from .fit import FitResult, fit, fit_detailed
|
|
33
|
+
from .repair import DEFAULT_PLACEHOLDER, repair, repair_detailed
|
|
34
|
+
from .tokens import approx_tokens
|
|
35
|
+
from .validate import validate
|
|
36
|
+
|
|
37
|
+
__version__ = "0.1.0"
|
|
38
|
+
|
|
39
|
+
__all__ = [
|
|
40
|
+
"DEFAULT_PLACEHOLDER",
|
|
41
|
+
"DUPLICATE_TOOL_CALL_ID",
|
|
42
|
+
"DUPLICATE_TOOL_RESULT",
|
|
43
|
+
"EMPTY_CONTENT",
|
|
44
|
+
"FIRST_MESSAGE_NOT_USER",
|
|
45
|
+
"MISPLACED_TOOL_RESULT",
|
|
46
|
+
"ORPHAN_TOOL_RESULT",
|
|
47
|
+
"SYSTEM_IN_MESSAGES",
|
|
48
|
+
"TOOL_RESULT_NOT_FIRST",
|
|
49
|
+
"UNANSWERED_TOOL_CALL",
|
|
50
|
+
"ContextBudgetError",
|
|
51
|
+
"FitResult",
|
|
52
|
+
"Fix",
|
|
53
|
+
"InvalidHistoryError",
|
|
54
|
+
"ProviderDetectionError",
|
|
55
|
+
"Report",
|
|
56
|
+
"TurnsafeError",
|
|
57
|
+
"Violation",
|
|
58
|
+
"__version__",
|
|
59
|
+
"approx_tokens",
|
|
60
|
+
"detect_provider",
|
|
61
|
+
"fit",
|
|
62
|
+
"fit_detailed",
|
|
63
|
+
"repair",
|
|
64
|
+
"repair_detailed",
|
|
65
|
+
"validate",
|
|
66
|
+
]
|
turnsafe/_match.py
ADDED
|
@@ -0,0 +1,126 @@
|
|
|
1
|
+
"""Pair every tool call with its tool result.
|
|
2
|
+
|
|
3
|
+
``validate`` reports from this and ``repair`` rebuilds from it, so the two can
|
|
4
|
+
never disagree about which result answers which call.
|
|
5
|
+
|
|
6
|
+
Matching rule: a call made by an assistant turn (one message, or for
|
|
7
|
+
Anthropic a run of consecutive assistant messages, which the API merges) is
|
|
8
|
+
answered by the first result with its id that comes after the turn and before
|
|
9
|
+
the next turn that calls the same id again. A result inside the provider's
|
|
10
|
+
allowed zone (see ``Adapter.zone``) is correctly placed; one elsewhere is
|
|
11
|
+
misplaced but recoverable. Further results for an answered call are
|
|
12
|
+
duplicates, and results that answer nothing are orphans.
|
|
13
|
+
|
|
14
|
+
Every result and call is visited a constant number of times, so matching is
|
|
15
|
+
linear in the size of the history.
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
from collections.abc import Sequence
|
|
21
|
+
from dataclasses import dataclass, field
|
|
22
|
+
|
|
23
|
+
from ._providers import Adapter
|
|
24
|
+
from ._types import Message
|
|
25
|
+
|
|
26
|
+
# A result item's location: (message index, position among that message's results).
|
|
27
|
+
Loc = tuple[int, int]
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@dataclass
|
|
31
|
+
class Match:
|
|
32
|
+
# assistant turn start -> call ids in order, duplicates removed
|
|
33
|
+
calls: dict[int, list[str]] = field(default_factory=dict)
|
|
34
|
+
# assistant turn start -> indices of the assistant messages in that turn
|
|
35
|
+
runs: dict[int, list[int]] = field(default_factory=dict)
|
|
36
|
+
# assistant index -> call id -> result location
|
|
37
|
+
answered: dict[int, dict[str, Loc]] = field(default_factory=dict)
|
|
38
|
+
# results that answer a call but sit outside the allowed zone
|
|
39
|
+
misplaced: set[Loc] = field(default_factory=set)
|
|
40
|
+
# assistant index -> call ids with no result anywhere in the window
|
|
41
|
+
missing: dict[int, list[str]] = field(default_factory=dict)
|
|
42
|
+
# (assistant index, call id) for ids repeated inside one message
|
|
43
|
+
duplicate_calls: list[tuple[int, str]] = field(default_factory=list)
|
|
44
|
+
# locations of second and later results for an already-answered call
|
|
45
|
+
duplicates: set[Loc] = field(default_factory=set)
|
|
46
|
+
# locations of results that answer nothing
|
|
47
|
+
orphans: set[Loc] = field(default_factory=set)
|
|
48
|
+
# result ids by location, for reporting
|
|
49
|
+
result_ids: dict[Loc, str] = field(default_factory=dict)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def match(adapter: Adapter, msgs: Sequence[Message]) -> Match:
|
|
53
|
+
m = Match()
|
|
54
|
+
results_by_id: dict[str, list[Loc]] = {}
|
|
55
|
+
for j, msg in enumerate(msgs):
|
|
56
|
+
for pos, (rid, _) in enumerate(adapter.result_items(msg)):
|
|
57
|
+
m.result_ids[(j, pos)] = rid
|
|
58
|
+
results_by_id.setdefault(rid, []).append((j, pos))
|
|
59
|
+
|
|
60
|
+
runs: list[list[int]] = []
|
|
61
|
+
for i, msg in enumerate(msgs):
|
|
62
|
+
if msg.get("role") != "assistant":
|
|
63
|
+
continue
|
|
64
|
+
if runs and runs[-1][-1] == i - 1 and adapter.merges_assistant_runs:
|
|
65
|
+
runs[-1].append(i)
|
|
66
|
+
else:
|
|
67
|
+
runs.append([i])
|
|
68
|
+
|
|
69
|
+
turn_ids: list[list[str]] = []
|
|
70
|
+
calls_by_id: dict[str, list[int]] = {}
|
|
71
|
+
for n, run in enumerate(runs):
|
|
72
|
+
ids: list[str] = []
|
|
73
|
+
seen: set[str] = set()
|
|
74
|
+
for k in run:
|
|
75
|
+
for cid in adapter.call_ids(msgs[k]):
|
|
76
|
+
if cid in seen:
|
|
77
|
+
m.duplicate_calls.append((run[0], cid))
|
|
78
|
+
else:
|
|
79
|
+
seen.add(cid)
|
|
80
|
+
ids.append(cid)
|
|
81
|
+
calls_by_id.setdefault(cid, []).append(n)
|
|
82
|
+
turn_ids.append(ids)
|
|
83
|
+
|
|
84
|
+
claimed: set[Loc] = set()
|
|
85
|
+
result_ptr: dict[str, int] = {}
|
|
86
|
+
call_ptr: dict[str, int] = {}
|
|
87
|
+
for n, run in enumerate(runs):
|
|
88
|
+
ids = turn_ids[n]
|
|
89
|
+
if not ids:
|
|
90
|
+
continue
|
|
91
|
+
i, end = run[0], run[-1]
|
|
92
|
+
m.calls[i] = ids
|
|
93
|
+
m.runs[i] = run
|
|
94
|
+
zone = set(adapter.zone(msgs, end))
|
|
95
|
+
answered: dict[str, Loc] = {}
|
|
96
|
+
for cid in ids:
|
|
97
|
+
# Stop at the next turn that reuses this id: its results are its own.
|
|
98
|
+
call_ptr[cid] = call_ptr.get(cid, 0) + 1
|
|
99
|
+
later = calls_by_id[cid][call_ptr[cid] :]
|
|
100
|
+
limit = runs[later[0]][0] if later else len(msgs)
|
|
101
|
+
|
|
102
|
+
found = results_by_id.get(cid, [])
|
|
103
|
+
k = result_ptr.get(cid, 0)
|
|
104
|
+
while k < len(found) and found[k][0] <= end:
|
|
105
|
+
k += 1 # before the call: left as orphans
|
|
106
|
+
locs = []
|
|
107
|
+
while k < len(found) and found[k][0] < limit:
|
|
108
|
+
locs.append(found[k])
|
|
109
|
+
k += 1
|
|
110
|
+
result_ptr[cid] = k
|
|
111
|
+
if not locs:
|
|
112
|
+
m.missing.setdefault(i, []).append(cid)
|
|
113
|
+
continue
|
|
114
|
+
# Prefer a correctly placed result over a misplaced one.
|
|
115
|
+
best = next((loc for loc in locs if loc[0] in zone), locs[0])
|
|
116
|
+
answered[cid] = best
|
|
117
|
+
if best[0] not in zone:
|
|
118
|
+
m.misplaced.add(best)
|
|
119
|
+
for loc in locs:
|
|
120
|
+
claimed.add(loc)
|
|
121
|
+
if loc != best:
|
|
122
|
+
m.duplicates.add(loc)
|
|
123
|
+
m.answered[i] = answered
|
|
124
|
+
|
|
125
|
+
m.orphans = {loc for loc in m.result_ids if loc not in claimed}
|
|
126
|
+
return m
|
turnsafe/_providers.py
ADDED
|
@@ -0,0 +1,351 @@
|
|
|
1
|
+
"""Per-provider knowledge: where tool calls and tool results live in a message.
|
|
2
|
+
|
|
3
|
+
Everything format-specific is here. ``validate``, ``repair`` and ``fit`` only
|
|
4
|
+
talk to the ``Adapter`` interface, so adding a provider means adding a class.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import copy
|
|
10
|
+
from collections.abc import Iterable, Mapping, Sequence
|
|
11
|
+
from typing import Any
|
|
12
|
+
|
|
13
|
+
from ._types import Message, Provider, ProviderDetectionError
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def to_dict(obj: Any, what: str = "message") -> dict[str, Any]:
|
|
17
|
+
"""Turn a message or block into a plain dict.
|
|
18
|
+
|
|
19
|
+
Accepts mappings and SDK objects with ``model_dump`` (openai / anthropic
|
|
20
|
+
pydantic models). ``None`` fields are dropped, which both APIs accept.
|
|
21
|
+
"""
|
|
22
|
+
if isinstance(obj, Mapping):
|
|
23
|
+
return dict(obj)
|
|
24
|
+
dump = getattr(obj, "model_dump", None)
|
|
25
|
+
if callable(dump):
|
|
26
|
+
result = dump(exclude_none=True)
|
|
27
|
+
if isinstance(result, dict):
|
|
28
|
+
return result
|
|
29
|
+
raise TypeError(
|
|
30
|
+
f"Expected a {what} dict (or an SDK object with model_dump()), got {type(obj).__name__}"
|
|
31
|
+
)
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def normalize(messages: Iterable[Any]) -> list[Message]:
|
|
35
|
+
"""Deep-copy the history into plain dicts so callers' objects are never mutated.
|
|
36
|
+
|
|
37
|
+
Tuples become lists, and tool call ids are checked: a call or result
|
|
38
|
+
without a usable id would silently escape every rule.
|
|
39
|
+
"""
|
|
40
|
+
if isinstance(messages, (str, bytes, Mapping)):
|
|
41
|
+
raise TypeError("messages must be a list of message dicts")
|
|
42
|
+
out: list[Message] = []
|
|
43
|
+
for i, raw in enumerate(messages):
|
|
44
|
+
msg = copy.deepcopy(to_dict(raw))
|
|
45
|
+
content = msg.get("content")
|
|
46
|
+
if isinstance(content, (list, tuple)):
|
|
47
|
+
msg["content"] = [to_dict(block, "content block") for block in content]
|
|
48
|
+
for k, block in enumerate(msg["content"]):
|
|
49
|
+
if block.get("type") in ("tool_use", "server_tool_use", "mcp_tool_use"):
|
|
50
|
+
_check_id(block.get("id"), f"messages[{i}].content[{k}].id")
|
|
51
|
+
elif block.get("type") == "tool_result":
|
|
52
|
+
_check_id(block.get("tool_use_id"), f"messages[{i}].content[{k}].tool_use_id")
|
|
53
|
+
calls = msg.get("tool_calls")
|
|
54
|
+
if isinstance(calls, (list, tuple)):
|
|
55
|
+
msg["tool_calls"] = [to_dict(call, "tool call") for call in calls]
|
|
56
|
+
for k, call in enumerate(msg["tool_calls"]):
|
|
57
|
+
_check_id(call.get("id"), f"messages[{i}].tool_calls[{k}].id")
|
|
58
|
+
if msg.get("role") == "tool":
|
|
59
|
+
_check_id(msg.get("tool_call_id"), f"messages[{i}].tool_call_id")
|
|
60
|
+
out.append(msg)
|
|
61
|
+
return out
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def _check_id(value: Any, where: str) -> None:
|
|
65
|
+
if not isinstance(value, str) or not value:
|
|
66
|
+
raise ValueError(f"{where} must be a non-empty string, got {value!r}")
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
class Adapter:
|
|
70
|
+
"""How one provider's message format stores tool calls and results.
|
|
71
|
+
|
|
72
|
+
A "result item" is whatever unit carries one tool result: a whole
|
|
73
|
+
``role: tool`` message for OpenAI, a ``tool_result`` block for Anthropic.
|
|
74
|
+
"""
|
|
75
|
+
|
|
76
|
+
name: Provider
|
|
77
|
+
# Whether the API treats consecutive assistant messages as one turn.
|
|
78
|
+
merges_assistant_runs = False
|
|
79
|
+
|
|
80
|
+
def call_ids(self, msg: Message) -> list[str]:
|
|
81
|
+
"""Ids of client tool calls made by this message (assistant only)."""
|
|
82
|
+
raise NotImplementedError
|
|
83
|
+
|
|
84
|
+
def result_items(self, msg: Message) -> list[tuple[str, Any]]:
|
|
85
|
+
"""(tool_call_id, item) for each tool result carried by this message."""
|
|
86
|
+
raise NotImplementedError
|
|
87
|
+
|
|
88
|
+
def zone(self, messages: Sequence[Message], i: int) -> list[int]:
|
|
89
|
+
"""Indices where results for the calls in ``messages[i]`` are allowed."""
|
|
90
|
+
raise NotImplementedError
|
|
91
|
+
|
|
92
|
+
def is_turn_start(self, msg: Message) -> bool:
|
|
93
|
+
"""True for a user message that starts a new turn (not just tool results)."""
|
|
94
|
+
raise NotImplementedError
|
|
95
|
+
|
|
96
|
+
def placeholder(self, tool_call_id: str, text: str) -> Any:
|
|
97
|
+
"""A synthetic result item for a call whose real result is missing."""
|
|
98
|
+
raise NotImplementedError
|
|
99
|
+
|
|
100
|
+
def strip_calls(self, msg: Message, ids: set[str]) -> Message | None:
|
|
101
|
+
"""``msg`` without the given calls, or None if nothing meaningful is left."""
|
|
102
|
+
raise NotImplementedError
|
|
103
|
+
|
|
104
|
+
def emit_group(
|
|
105
|
+
self, assistants: list[Message], items: list[Any], zone_msgs: list[Message]
|
|
106
|
+
) -> list[Message]:
|
|
107
|
+
"""The assistant turn followed by messages that carry ``items``.
|
|
108
|
+
|
|
109
|
+
``zone_msgs`` are the input messages in the assistant's zone. Their
|
|
110
|
+
result items are replaced by ``items``; everything else they hold,
|
|
111
|
+
content and keys, must be kept.
|
|
112
|
+
"""
|
|
113
|
+
raise NotImplementedError
|
|
114
|
+
|
|
115
|
+
def rename_call(self, msg: Message, old: str, new: str) -> Message:
|
|
116
|
+
"""``msg`` with the tool call ``old`` renamed to ``new``."""
|
|
117
|
+
raise NotImplementedError
|
|
118
|
+
|
|
119
|
+
def rename_result(self, item: Any, new: str) -> Any:
|
|
120
|
+
"""A copy of result ``item`` answering ``new`` instead."""
|
|
121
|
+
raise NotImplementedError
|
|
122
|
+
|
|
123
|
+
def without_results(self, msg: Message) -> Message | None:
|
|
124
|
+
"""``msg`` with every tool result removed, or None if it becomes empty."""
|
|
125
|
+
raise NotImplementedError
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
class OpenAIAdapter(Adapter):
|
|
129
|
+
"""OpenAI Chat Completions (``role: tool`` messages answer ``tool_calls``)."""
|
|
130
|
+
|
|
131
|
+
name: Provider = "openai"
|
|
132
|
+
|
|
133
|
+
def call_ids(self, msg: Message) -> list[str]:
|
|
134
|
+
if msg.get("role") != "assistant":
|
|
135
|
+
return []
|
|
136
|
+
calls = msg.get("tool_calls") or []
|
|
137
|
+
return [c["id"] for c in calls if isinstance(c, dict) and isinstance(c.get("id"), str)]
|
|
138
|
+
|
|
139
|
+
def result_items(self, msg: Message) -> list[tuple[str, Any]]:
|
|
140
|
+
if msg.get("role") == "tool":
|
|
141
|
+
return [(str(msg.get("tool_call_id", "")), msg)]
|
|
142
|
+
return []
|
|
143
|
+
|
|
144
|
+
def zone(self, messages: Sequence[Message], i: int) -> list[int]:
|
|
145
|
+
out = []
|
|
146
|
+
j = i + 1
|
|
147
|
+
while j < len(messages) and messages[j].get("role") == "tool":
|
|
148
|
+
out.append(j)
|
|
149
|
+
j += 1
|
|
150
|
+
return out
|
|
151
|
+
|
|
152
|
+
def is_turn_start(self, msg: Message) -> bool:
|
|
153
|
+
return msg.get("role") == "user"
|
|
154
|
+
|
|
155
|
+
def placeholder(self, tool_call_id: str, text: str) -> Any:
|
|
156
|
+
return {"role": "tool", "tool_call_id": tool_call_id, "content": text}
|
|
157
|
+
|
|
158
|
+
def strip_calls(self, msg: Message, ids: set[str]) -> Message | None:
|
|
159
|
+
out = dict(msg)
|
|
160
|
+
kept = [c for c in msg.get("tool_calls") or [] if c.get("id") not in ids]
|
|
161
|
+
if kept:
|
|
162
|
+
out["tool_calls"] = kept
|
|
163
|
+
return out
|
|
164
|
+
out.pop("tool_calls", None)
|
|
165
|
+
out.pop("parallel_tool_calls", None)
|
|
166
|
+
if _is_empty(out.get("content")) and not out.get("audio") and not out.get("refusal"):
|
|
167
|
+
return None
|
|
168
|
+
return out
|
|
169
|
+
|
|
170
|
+
def emit_group(
|
|
171
|
+
self, assistants: list[Message], items: list[Any], zone_msgs: list[Message]
|
|
172
|
+
) -> list[Message]:
|
|
173
|
+
# The zone only ever holds tool messages, which are all in ``items`` or dropped.
|
|
174
|
+
return [*assistants, *items]
|
|
175
|
+
|
|
176
|
+
def rename_call(self, msg: Message, old: str, new: str) -> Message:
|
|
177
|
+
calls = [{**c, "id": new} if c.get("id") == old else c for c in msg["tool_calls"]]
|
|
178
|
+
return {**msg, "tool_calls": calls}
|
|
179
|
+
|
|
180
|
+
def rename_result(self, item: Any, new: str) -> Any:
|
|
181
|
+
return {**item, "tool_call_id": new}
|
|
182
|
+
|
|
183
|
+
def without_results(self, msg: Message) -> Message | None:
|
|
184
|
+
return None if msg.get("role") == "tool" else msg
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
# Anthropic blocks that pair with a later client-side tool_result. Server tool
|
|
188
|
+
# blocks (server_tool_use, web_search_tool_result, mcp_tool_use, ...) are
|
|
189
|
+
# answered inside the same assistant message, so they are left alone.
|
|
190
|
+
_ANTHROPIC_CALL = "tool_use"
|
|
191
|
+
_ANTHROPIC_RESULT = "tool_result"
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
def _blocks(msg: Message) -> list[Any]:
|
|
195
|
+
content = msg.get("content")
|
|
196
|
+
return content if isinstance(content, list) else []
|
|
197
|
+
|
|
198
|
+
|
|
199
|
+
class AnthropicAdapter(Adapter):
|
|
200
|
+
"""Anthropic Messages (``tool_result`` blocks answer ``tool_use`` blocks)."""
|
|
201
|
+
|
|
202
|
+
name: Provider = "anthropic"
|
|
203
|
+
# "Consecutive user or assistant turns in your request will be combined
|
|
204
|
+
# into a single turn." https://platform.claude.com/docs/en/api/messages
|
|
205
|
+
merges_assistant_runs = True
|
|
206
|
+
|
|
207
|
+
def call_ids(self, msg: Message) -> list[str]:
|
|
208
|
+
if msg.get("role") != "assistant":
|
|
209
|
+
return []
|
|
210
|
+
return [
|
|
211
|
+
b["id"]
|
|
212
|
+
for b in _blocks(msg)
|
|
213
|
+
if isinstance(b, dict)
|
|
214
|
+
and b.get("type") == _ANTHROPIC_CALL
|
|
215
|
+
and isinstance(b.get("id"), str)
|
|
216
|
+
]
|
|
217
|
+
|
|
218
|
+
def result_items(self, msg: Message) -> list[tuple[str, Any]]:
|
|
219
|
+
if msg.get("role") != "user":
|
|
220
|
+
return []
|
|
221
|
+
return [
|
|
222
|
+
(str(b.get("tool_use_id", "")), b)
|
|
223
|
+
for b in _blocks(msg)
|
|
224
|
+
if isinstance(b, dict) and b.get("type") == _ANTHROPIC_RESULT
|
|
225
|
+
]
|
|
226
|
+
|
|
227
|
+
def zone(self, messages: Sequence[Message], i: int) -> list[int]:
|
|
228
|
+
# The API merges consecutive user messages into one turn before
|
|
229
|
+
# checking, so results may be spread over the whole run.
|
|
230
|
+
out = []
|
|
231
|
+
j = i + 1
|
|
232
|
+
while j < len(messages) and messages[j].get("role") == "user":
|
|
233
|
+
out.append(j)
|
|
234
|
+
j += 1
|
|
235
|
+
return out
|
|
236
|
+
|
|
237
|
+
def is_turn_start(self, msg: Message) -> bool:
|
|
238
|
+
return msg.get("role") == "user" and not self.result_items(msg)
|
|
239
|
+
|
|
240
|
+
def placeholder(self, tool_call_id: str, text: str) -> Any:
|
|
241
|
+
return {
|
|
242
|
+
"type": "tool_result",
|
|
243
|
+
"tool_use_id": tool_call_id,
|
|
244
|
+
"content": text,
|
|
245
|
+
"is_error": True,
|
|
246
|
+
}
|
|
247
|
+
|
|
248
|
+
def strip_calls(self, msg: Message, ids: set[str]) -> Message | None:
|
|
249
|
+
kept = [
|
|
250
|
+
b for b in _blocks(msg) if not (b.get("type") == _ANTHROPIC_CALL and b.get("id") in ids)
|
|
251
|
+
]
|
|
252
|
+
# Thinking blocks alone are not a meaningful assistant turn.
|
|
253
|
+
if not any(b.get("type") not in ("thinking", "redacted_thinking") for b in kept):
|
|
254
|
+
return None
|
|
255
|
+
return {**msg, "content": kept}
|
|
256
|
+
|
|
257
|
+
def emit_group(
|
|
258
|
+
self, assistants: list[Message], items: list[Any], zone_msgs: list[Message]
|
|
259
|
+
) -> list[Message]:
|
|
260
|
+
# One user message: every tool_result first, then whatever else the zone
|
|
261
|
+
# held, in order. Other keys from every zone message are kept (first wins).
|
|
262
|
+
if not items:
|
|
263
|
+
kept = [self.without_results(z) for z in zone_msgs]
|
|
264
|
+
return [*assistants, *(z for z in kept if z is not None)]
|
|
265
|
+
base: Message = {}
|
|
266
|
+
rest: list[Any] = []
|
|
267
|
+
for msg in zone_msgs:
|
|
268
|
+
for key, value in msg.items():
|
|
269
|
+
if key != "content":
|
|
270
|
+
base.setdefault(key, value)
|
|
271
|
+
content = msg.get("content")
|
|
272
|
+
if isinstance(content, str):
|
|
273
|
+
rest += [{"type": "text", "text": content}] if content else []
|
|
274
|
+
elif isinstance(content, list):
|
|
275
|
+
rest += [b for b in content if b.get("type") != _ANTHROPIC_RESULT]
|
|
276
|
+
base.setdefault("role", "user")
|
|
277
|
+
return [*assistants, {**base, "content": [*items, *rest]}]
|
|
278
|
+
|
|
279
|
+
def rename_call(self, msg: Message, old: str, new: str) -> Message:
|
|
280
|
+
content = [
|
|
281
|
+
{**b, "id": new} if b.get("type") == _ANTHROPIC_CALL and b.get("id") == old else b
|
|
282
|
+
for b in _blocks(msg)
|
|
283
|
+
]
|
|
284
|
+
return {**msg, "content": content}
|
|
285
|
+
|
|
286
|
+
def rename_result(self, item: Any, new: str) -> Any:
|
|
287
|
+
return {**item, "tool_use_id": new}
|
|
288
|
+
|
|
289
|
+
def without_results(self, msg: Message) -> Message | None:
|
|
290
|
+
if not self.result_items(msg):
|
|
291
|
+
return msg
|
|
292
|
+
kept = [b for b in _blocks(msg) if b.get("type") != _ANTHROPIC_RESULT]
|
|
293
|
+
if not kept:
|
|
294
|
+
return None
|
|
295
|
+
return {**msg, "content": kept}
|
|
296
|
+
|
|
297
|
+
|
|
298
|
+
def _is_empty(content: Any) -> bool:
|
|
299
|
+
return content is None or content == "" or content == []
|
|
300
|
+
|
|
301
|
+
|
|
302
|
+
ADAPTERS: dict[str, Adapter] = {"openai": OpenAIAdapter(), "anthropic": AnthropicAdapter()}
|
|
303
|
+
|
|
304
|
+
_OPENAI_ROLES = {"tool", "developer", "function"}
|
|
305
|
+
_ANTHROPIC_BLOCKS = {
|
|
306
|
+
"tool_use",
|
|
307
|
+
"tool_result",
|
|
308
|
+
"thinking",
|
|
309
|
+
"redacted_thinking",
|
|
310
|
+
"server_tool_use",
|
|
311
|
+
"web_search_tool_result",
|
|
312
|
+
"mcp_tool_use",
|
|
313
|
+
"mcp_tool_result",
|
|
314
|
+
}
|
|
315
|
+
|
|
316
|
+
|
|
317
|
+
def detect(messages: Sequence[Message]) -> Provider:
|
|
318
|
+
"""Guess the message format from tool-call markers.
|
|
319
|
+
|
|
320
|
+
Plain text histories with no markers are reported as ``"openai"``; the
|
|
321
|
+
two formats validate identically when there are no tools involved, except
|
|
322
|
+
for Anthropic's content rules, so pass ``provider=`` when you know it.
|
|
323
|
+
"""
|
|
324
|
+
# "system" is deliberately not an OpenAI signal: in an Anthropic history it
|
|
325
|
+
# is a mistake that validate() should report, not a reason to give up.
|
|
326
|
+
openai = anthropic = False
|
|
327
|
+
for msg in messages:
|
|
328
|
+
if (
|
|
329
|
+
msg.get("role") in _OPENAI_ROLES
|
|
330
|
+
or msg.get("tool_calls")
|
|
331
|
+
or msg.get("tool_call_id") is not None
|
|
332
|
+
):
|
|
333
|
+
openai = True
|
|
334
|
+
if any(isinstance(b, dict) and b.get("type") in _ANTHROPIC_BLOCKS for b in _blocks(msg)):
|
|
335
|
+
anthropic = True
|
|
336
|
+
if openai and anthropic:
|
|
337
|
+
raise ProviderDetectionError(
|
|
338
|
+
"History mixes OpenAI and Anthropic tool formats; pass provider= explicitly."
|
|
339
|
+
)
|
|
340
|
+
return "anthropic" if anthropic else "openai"
|
|
341
|
+
|
|
342
|
+
|
|
343
|
+
def get_adapter(provider: str, messages: Sequence[Message]) -> Adapter:
|
|
344
|
+
if provider == "auto":
|
|
345
|
+
provider = detect(messages)
|
|
346
|
+
try:
|
|
347
|
+
return ADAPTERS[provider]
|
|
348
|
+
except KeyError:
|
|
349
|
+
raise ValueError(
|
|
350
|
+
f"Unknown provider {provider!r}; expected 'openai', 'anthropic' or 'auto'."
|
|
351
|
+
) from None
|
turnsafe/_thinking.py
ADDED
|
@@ -0,0 +1,99 @@
|
|
|
1
|
+
"""Keep Anthropic thinking blocks valid after the history has been changed.
|
|
2
|
+
|
|
3
|
+
On models with thinking-block prefix binding (Claude Fable 5.1 and Opus 5.5,
|
|
4
|
+
enforced by default for accounts created on or after 2026-08-31), a replayed
|
|
5
|
+
``thinking`` or ``redacted_thinking`` block is only valid while every message
|
|
6
|
+
before it is unchanged. Trimming from the front, dropping a message or
|
|
7
|
+
editing one invalidates that block and every later one, and the API answers
|
|
8
|
+
"Invalid `signature` in `thinking` block. The block is bound to a different
|
|
9
|
+
conversation."
|
|
10
|
+
|
|
11
|
+
The documented recovery is to remove the thinking blocks from the first
|
|
12
|
+
changed position onward, leaving everything else in place, or to set
|
|
13
|
+
``prefix_mismatch_behavior: "drop_block"`` and let the API drop them.
|
|
14
|
+
https://platform.claude.com/docs/en/build-with-claude/thinking-troubleshooting
|
|
15
|
+
|
|
16
|
+
Models without prefix binding differ: they want the latest assistant turn's
|
|
17
|
+
thinking replayed unchanged, and manual extended thinking requires it in a
|
|
18
|
+
tool loop. That is why ``"keep"`` is the default and stripping is opt-in.
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
from __future__ import annotations
|
|
22
|
+
|
|
23
|
+
from collections.abc import Sequence
|
|
24
|
+
from typing import Literal
|
|
25
|
+
|
|
26
|
+
from ._types import Message
|
|
27
|
+
|
|
28
|
+
ThinkingPolicy = Literal["strip_after_change", "keep"]
|
|
29
|
+
THINKING_BLOCKS = ("thinking", "redacted_thinking")
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def first_change(before: Sequence[Message], after: Sequence[Message]) -> int:
|
|
33
|
+
"""Index of the first message in ``after`` that differs from ``before``."""
|
|
34
|
+
for n, (a, b) in enumerate(zip(before, after, strict=False)):
|
|
35
|
+
if a != b:
|
|
36
|
+
return n
|
|
37
|
+
return min(len(before), len(after))
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def strip_after_change(
|
|
41
|
+
before: Sequence[Message], after: list[Message], src: Sequence[int] | None = None
|
|
42
|
+
) -> tuple[list[Message], list[int]]:
|
|
43
|
+
"""Remove every thinking block in ``after`` that follows its first change.
|
|
44
|
+
|
|
45
|
+
``src[n]`` is the index in ``before`` that ``after[n]`` was derived from.
|
|
46
|
+
When the first changed message is an edited version of the message at the
|
|
47
|
+
same position, thinking blocks that come before its first changed block are
|
|
48
|
+
kept: nothing before them changed.
|
|
49
|
+
"""
|
|
50
|
+
n = first_change(before, after)
|
|
51
|
+
if n >= len(after):
|
|
52
|
+
return after, []
|
|
53
|
+
msg = after[n]
|
|
54
|
+
content = msg.get("content")
|
|
55
|
+
same = src is not None and n < len(src) and src[n] == n and n < len(before)
|
|
56
|
+
old = before[n].get("content") if same else None
|
|
57
|
+
if msg.get("role") == "assistant" and isinstance(content, list) and isinstance(old, list):
|
|
58
|
+
b = next(
|
|
59
|
+
(k for k, (x, y) in enumerate(zip(old, content, strict=False)) if x != y),
|
|
60
|
+
min(len(old), len(content)),
|
|
61
|
+
)
|
|
62
|
+
kept = content[:b] + [x for x in content[b:] if x.get("type") not in THINKING_BLOCKS]
|
|
63
|
+
head: list[Message] = after[:n]
|
|
64
|
+
touched = [n] if len(kept) != len(content) else []
|
|
65
|
+
if kept:
|
|
66
|
+
head.append({**msg, "content": kept})
|
|
67
|
+
rest, more = strip_from(after[n + 1 :], 0)
|
|
68
|
+
return head + rest, touched + [n + 1 + k for k in more]
|
|
69
|
+
return strip_from(after, n)
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def strip_from(msgs: list[Message], start: int) -> tuple[list[Message], list[int]]:
|
|
73
|
+
"""Remove thinking blocks from assistant messages at ``start`` and later.
|
|
74
|
+
|
|
75
|
+
Returns the new list and the output positions of the messages that lost
|
|
76
|
+
blocks. An assistant message left with nothing but thinking is dropped;
|
|
77
|
+
its position is reported too.
|
|
78
|
+
"""
|
|
79
|
+
out: list[Message] = msgs[:start]
|
|
80
|
+
touched: list[int] = []
|
|
81
|
+
for n in range(start, len(msgs)):
|
|
82
|
+
msg = msgs[n]
|
|
83
|
+
content = msg.get("content")
|
|
84
|
+
if msg.get("role") != "assistant" or not isinstance(content, list):
|
|
85
|
+
out.append(msg)
|
|
86
|
+
continue
|
|
87
|
+
kept = [b for b in content if b.get("type") not in THINKING_BLOCKS]
|
|
88
|
+
if len(kept) == len(content):
|
|
89
|
+
out.append(msg)
|
|
90
|
+
continue
|
|
91
|
+
touched.append(n)
|
|
92
|
+
if kept:
|
|
93
|
+
out.append({**msg, "content": kept})
|
|
94
|
+
return out, touched
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def check_policy(thinking: str) -> None:
|
|
98
|
+
if thinking not in ("strip_after_change", "keep"):
|
|
99
|
+
raise ValueError(f"thinking must be 'strip_after_change' or 'keep', got {thinking!r}")
|