recover concatenated artifact writes

This commit is contained in:
pewdiepie-archdaemon
2026-09-19 01:53:12 +00:00
parent b32fa31b37
commit 2308ff9b80
4 changed files with 212 additions and 5 deletions
+88 -3
View File
@@ -624,11 +624,22 @@ def explicit_text_artifact_target(user_text):
def malformed_write_handoff_target(arguments, required_artifacts=(), user_text=''):
"""Recover one textual write target even when no prompt path was parsed."""
candidates = tuple(required_artifacts or ())
target = candidates[0] if len(candidates) == 1 else (
_partial_json_string_field(arguments, 'path')
or explicit_text_artifact_target(user_text)
recovered = _partial_json_string_field(arguments, 'path')
sole_required = str(candidates[0] or '').strip() if len(candidates) == 1 else ''
# A runner may require a directory containing several outputs. It is not
# itself a writable file target, so prefer the concrete descendant path in
# the malformed call. Exact required files remain authoritative.
target = (
sole_required
if sole_required and Path(sole_required).suffix
else recovered or explicit_text_artifact_target(user_text)
)
target = str(target or '').strip()
if sole_required and not Path(sole_required).suffix:
required_root = os.path.normpath(sole_required)
normalized_target = os.path.normpath(target) if target else ''
if not normalized_target.startswith(required_root + os.sep):
return ''
if not target or Path(target).suffix.lower() in _NON_TEXT_ARTIFACT_SUFFIXES:
return ''
return target
@@ -1637,6 +1648,72 @@ def protocol_safe_tool_calls(calls):
return safe_calls
def expand_concatenated_write_calls(calls, *, max_calls=16):
"""Split one unambiguous stream of write_file JSON objects.
Some OpenAI-compatible providers serialize a requested multi-call batch as
adjacent JSON objects inside one ``arguments`` string. Recover only the
narrow write-file shape; all expanded calls still pass the ordinary schema,
policy, workspace, repeat, and execution checks later in the loop.
"""
expanded = []
recovered_count = 0
decoder = json.JSONDecoder()
for call in calls or ():
function = call.get('function') or {}
raw = function.get('arguments') or ''
if canonical(function.get('name', '')) != 'write_file':
expanded.append(call)
continue
try:
json.loads(raw)
except (TypeError, ValueError, json.JSONDecodeError):
pass
else:
expanded.append(call)
continue
if not isinstance(raw, str):
expanded.append(call)
continue
items = []
cursor = 0
try:
while cursor < len(raw):
while cursor < len(raw) and raw[cursor].isspace():
cursor += 1
if cursor >= len(raw):
break
item, cursor = decoder.raw_decode(raw, cursor)
items.append(item)
if len(items) > max_calls:
raise ValueError('too many concatenated write calls')
except (TypeError, ValueError, json.JSONDecodeError):
expanded.append(call)
continue
if not (2 <= len(items) <= max_calls) or not all(
isinstance(item, dict)
and set(item) == {'path', 'content'}
and isinstance(item['path'], str)
and bool(item['path'].strip())
and isinstance(item['content'], str)
for item in items
):
expanded.append(call)
continue
base_id = str(call.get('id') or 'call_write')
for index, item in enumerate(items):
expanded.append({
'id': f'{base_id}_{index}',
'type': 'function',
'function': {
'name': function.get('name', 'write_file'),
'arguments': json.dumps(item, ensure_ascii=False),
},
})
recovered_count += len(items)
return expanded, recovered_count
def artifact_body_from_handoff(response):
"""Extract a complete textual artifact body from a no-tools recovery turn."""
raw = str(response or '').strip()
@@ -4746,6 +4823,14 @@ async def stream_preview(*, endpoint_url, model, messages, headers, turn_contrac
})
else:
unexecutable_dsml_completion = True
proposed, recovered_write_calls = expand_concatenated_write_calls(proposed)
if recovered_write_calls:
yield event({
'type': 'tool_argument_recovery',
'format': 'concatenated_json_objects',
'round': round_number,
'calls': recovered_write_calls,
})
proposed = serialize_required_email_attachment_chain(
proposed, contract_required_tools, executions,
)