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 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}")