add test guide cases
This commit is contained in:
+24
-1
@@ -26,9 +26,13 @@ class ManagedAgentGroup:
|
||||
label: str | None = None
|
||||
parent_agent_id: str | None = None
|
||||
child_agent_ids: tuple[str, ...] = ()
|
||||
strategy: str = 'serial'
|
||||
status: str = 'running'
|
||||
completed_children: int = 0
|
||||
failed_children: int = 0
|
||||
batch_count: int = 0
|
||||
max_batch_size: int = 0
|
||||
dependency_skips: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -68,6 +72,7 @@ class AgentManager:
|
||||
*,
|
||||
label: str | None = None,
|
||||
parent_agent_id: str | None = None,
|
||||
strategy: str = 'serial',
|
||||
) -> str:
|
||||
self._group_counter += 1
|
||||
group_id = f'group_{self._group_counter}'
|
||||
@@ -75,6 +80,7 @@ class AgentManager:
|
||||
group_id=group_id,
|
||||
label=label,
|
||||
parent_agent_id=parent_agent_id,
|
||||
strategy=strategy,
|
||||
)
|
||||
return group_id
|
||||
|
||||
@@ -97,9 +103,13 @@ class AgentManager:
|
||||
label=group.label,
|
||||
parent_agent_id=group.parent_agent_id,
|
||||
child_agent_ids=updated_children,
|
||||
strategy=group.strategy,
|
||||
status=group.status,
|
||||
completed_children=group.completed_children,
|
||||
failed_children=group.failed_children,
|
||||
batch_count=group.batch_count,
|
||||
max_batch_size=group.max_batch_size,
|
||||
dependency_skips=group.dependency_skips,
|
||||
)
|
||||
record = self.records.get(agent_id)
|
||||
if record is None:
|
||||
@@ -129,6 +139,9 @@ class AgentManager:
|
||||
status: str,
|
||||
completed_children: int,
|
||||
failed_children: int,
|
||||
batch_count: int = 0,
|
||||
max_batch_size: int = 0,
|
||||
dependency_skips: int = 0,
|
||||
) -> None:
|
||||
group = self.groups.get(group_id)
|
||||
if group is None:
|
||||
@@ -138,9 +151,13 @@ class AgentManager:
|
||||
label=group.label,
|
||||
parent_agent_id=group.parent_agent_id,
|
||||
child_agent_ids=group.child_agent_ids,
|
||||
strategy=group.strategy,
|
||||
status=status,
|
||||
completed_children=completed_children,
|
||||
failed_children=failed_children,
|
||||
batch_count=batch_count,
|
||||
max_batch_size=max_batch_size,
|
||||
dependency_skips=dependency_skips,
|
||||
)
|
||||
|
||||
def finish_agent(
|
||||
@@ -209,11 +226,15 @@ class AgentManager:
|
||||
return {
|
||||
'group_id': group.group_id,
|
||||
'label': group.label,
|
||||
'strategy': group.strategy,
|
||||
'status': group.status,
|
||||
'child_count': len(children),
|
||||
'completed_children': group.completed_children,
|
||||
'failed_children': group.failed_children,
|
||||
'resumed_children': resumed_children,
|
||||
'batch_count': group.batch_count,
|
||||
'max_batch_size': group.max_batch_size,
|
||||
'dependency_skips': group.dependency_skips,
|
||||
'stop_reason_counts': stop_reason_counts,
|
||||
}
|
||||
|
||||
@@ -266,7 +287,9 @@ class AgentManager:
|
||||
lines.append(
|
||||
f'- {label}: group_status={group.status} children={len(group.child_agent_ids)} '
|
||||
f'completed={group.completed_children} failed={group.failed_children} '
|
||||
f"resumed={summary['resumed_children']}{stop_suffix}"
|
||||
f"resumed={summary['resumed_children']} strategy={group.strategy} "
|
||||
f"batches={group.batch_count} max_batch_size={group.max_batch_size} "
|
||||
f"dependency_skips={group.dependency_skips}{stop_suffix}"
|
||||
)
|
||||
if len(self.groups) > 6:
|
||||
lines.append(f'- ... plus {len(self.groups) - 6} more agent groups')
|
||||
|
||||
+649
-92
File diff suppressed because it is too large
Load Diff
@@ -92,6 +92,7 @@ class AgentSessionState:
|
||||
user_context: dict[str, str] = field(default_factory=dict)
|
||||
system_context: dict[str, str] = field(default_factory=dict)
|
||||
messages: list[AgentMessage] = field(default_factory=list)
|
||||
mutation_serial: int = 0
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
@@ -200,6 +201,7 @@ class AgentSessionState:
|
||||
previous_content=message.content,
|
||||
previous_state=message.state,
|
||||
previous_stop_reason=message.stop_reason,
|
||||
mutation_serial=self._next_mutation_serial(),
|
||||
)
|
||||
merged_metadata = _advance_lineage_revision(merged_metadata)
|
||||
self.messages[index] = replace(
|
||||
@@ -246,6 +248,7 @@ class AgentSessionState:
|
||||
previous_content=message.content,
|
||||
previous_state=message.state,
|
||||
previous_stop_reason=message.stop_reason,
|
||||
mutation_serial=self._next_mutation_serial(),
|
||||
)
|
||||
merged_metadata = _advance_lineage_revision(merged_metadata)
|
||||
self.messages[index] = replace(
|
||||
@@ -269,6 +272,7 @@ class AgentSessionState:
|
||||
previous_content=message.content,
|
||||
previous_state=message.state,
|
||||
previous_stop_reason=message.stop_reason,
|
||||
mutation_serial=self._next_mutation_serial(),
|
||||
)
|
||||
merged_metadata = _advance_lineage_revision(merged_metadata)
|
||||
self.messages[index] = replace(
|
||||
@@ -362,6 +366,7 @@ class AgentSessionState:
|
||||
previous_content=message.content,
|
||||
previous_state=message.state,
|
||||
previous_stop_reason=message.stop_reason,
|
||||
mutation_serial=self._next_mutation_serial(),
|
||||
)
|
||||
merged_metadata = _advance_lineage_revision(merged_metadata)
|
||||
if metadata:
|
||||
@@ -391,6 +396,7 @@ class AgentSessionState:
|
||||
previous_content=message.content,
|
||||
previous_state=message.state,
|
||||
previous_stop_reason=message.stop_reason,
|
||||
mutation_serial=self._next_mutation_serial(),
|
||||
)
|
||||
merged_metadata = _advance_lineage_revision(merged_metadata)
|
||||
if metadata:
|
||||
@@ -430,6 +436,7 @@ class AgentSessionState:
|
||||
previous_content=message.content,
|
||||
previous_state=message.state,
|
||||
previous_stop_reason=message.stop_reason,
|
||||
mutation_serial=self._next_mutation_serial(),
|
||||
)
|
||||
merged_metadata = _advance_lineage_revision(merged_metadata)
|
||||
if metadata:
|
||||
@@ -474,6 +481,10 @@ class AgentSessionState:
|
||||
def transcript(self) -> tuple[JSONDict, ...]:
|
||||
return tuple(message.to_transcript_entry() for message in self.messages)
|
||||
|
||||
def _next_mutation_serial(self) -> int:
|
||||
self.mutation_serial += 1
|
||||
return self.mutation_serial
|
||||
|
||||
@classmethod
|
||||
def from_persisted(
|
||||
cls,
|
||||
@@ -488,6 +499,17 @@ class AgentSessionState:
|
||||
user_context=dict(user_context or {}),
|
||||
system_context=dict(system_context or {}),
|
||||
messages=[AgentMessage.from_openai_message(message) for message in messages],
|
||||
mutation_serial=max(
|
||||
(
|
||||
int(message.get('metadata', {}).get('last_mutation_serial', 0))
|
||||
for message in messages
|
||||
if isinstance(message, dict)
|
||||
and isinstance(message.get('metadata'), dict)
|
||||
and isinstance(message.get('metadata', {}).get('last_mutation_serial', 0), int)
|
||||
and not isinstance(message.get('metadata', {}).get('last_mutation_serial', 0), bool)
|
||||
),
|
||||
default=0,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -522,6 +544,7 @@ def _record_mutation(
|
||||
previous_content: str,
|
||||
previous_state: str,
|
||||
previous_stop_reason: str | None,
|
||||
mutation_serial: int,
|
||||
) -> JSONDict:
|
||||
mutations = metadata.get('mutations')
|
||||
if not isinstance(mutations, list):
|
||||
@@ -538,6 +561,7 @@ def _record_mutation(
|
||||
'previous_stop_reason': previous_stop_reason,
|
||||
'previous_content_length': len(previous_content),
|
||||
'previous_content_preview': preview or '(empty)',
|
||||
'serial': mutation_serial,
|
||||
}
|
||||
)
|
||||
if len(mutations) > MAX_MUTATION_HISTORY:
|
||||
@@ -545,6 +569,11 @@ def _record_mutation(
|
||||
metadata['mutations'] = mutations
|
||||
metadata['mutation_count'] = len(mutations)
|
||||
metadata['last_mutation_kind'] = mutation_kind
|
||||
metadata['last_mutation_serial'] = mutation_serial
|
||||
max_mutation_serial = metadata.get('max_mutation_serial')
|
||||
if isinstance(max_mutation_serial, bool) or not isinstance(max_mutation_serial, int):
|
||||
max_mutation_serial = 0
|
||||
metadata['max_mutation_serial'] = max(max_mutation_serial, mutation_serial)
|
||||
totals = metadata.get('mutation_totals')
|
||||
if not isinstance(totals, dict):
|
||||
totals = {}
|
||||
|
||||
@@ -242,6 +242,10 @@ def default_tool_registry() -> dict[str, AgentTool]:
|
||||
'max_turns': {'type': 'integer', 'minimum': 1, 'maximum': 20},
|
||||
'resume_session_id': {'type': 'string'},
|
||||
'session_id': {'type': 'string'},
|
||||
'depends_on': {
|
||||
'type': 'array',
|
||||
'items': {'type': 'string'},
|
||||
},
|
||||
},
|
||||
'required': ['prompt'],
|
||||
},
|
||||
@@ -255,6 +259,8 @@ def default_tool_registry() -> dict[str, AgentTool]:
|
||||
'allow_shell': {'type': 'boolean'},
|
||||
'include_parent_context': {'type': 'boolean'},
|
||||
'continue_on_error': {'type': 'boolean'},
|
||||
'max_failures': {'type': 'integer', 'minimum': 0, 'maximum': 20},
|
||||
'strategy': {'type': 'string'},
|
||||
},
|
||||
},
|
||||
handler=_delegate_agent_placeholder,
|
||||
|
||||
@@ -80,6 +80,8 @@ class BudgetConfig:
|
||||
max_total_cost_usd: float | None = None
|
||||
max_tool_calls: int | None = None
|
||||
max_delegated_tasks: int | None = None
|
||||
max_model_calls: int | None = None
|
||||
max_session_turns: int | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
||||
+85
@@ -5,6 +5,7 @@ import os
|
||||
from pathlib import Path
|
||||
from dataclasses import replace
|
||||
import json
|
||||
from typing import Callable
|
||||
|
||||
from .agent_runtime import LocalCodingAgent
|
||||
from .agent_types import (
|
||||
@@ -63,6 +64,8 @@ def _add_agent_common_args(parser: argparse.ArgumentParser, *, include_backend:
|
||||
parser.add_argument('--max-budget-usd', type=float)
|
||||
parser.add_argument('--max-tool-calls', type=int)
|
||||
parser.add_argument('--max-delegated-tasks', type=int)
|
||||
parser.add_argument('--max-model-calls', type=int)
|
||||
parser.add_argument('--max-session-turns', type=int)
|
||||
parser.add_argument('--response-schema-file')
|
||||
parser.add_argument('--response-schema-name')
|
||||
parser.add_argument('--response-schema-strict', action='store_true')
|
||||
@@ -95,6 +98,8 @@ def _build_runtime_config(args: argparse.Namespace) -> AgentRuntimeConfig:
|
||||
max_total_cost_usd=getattr(args, 'max_budget_usd', None),
|
||||
max_tool_calls=getattr(args, 'max_tool_calls', None),
|
||||
max_delegated_tasks=getattr(args, 'max_delegated_tasks', None),
|
||||
max_model_calls=getattr(args, 'max_model_calls', None),
|
||||
max_session_turns=getattr(args, 'max_session_turns', None),
|
||||
),
|
||||
output_schema=_load_output_schema_config(args),
|
||||
session_directory=(Path('.port_sessions') / 'agent').resolve(),
|
||||
@@ -175,6 +180,8 @@ def _add_agent_resume_args(parser: argparse.ArgumentParser) -> None:
|
||||
parser.add_argument('--max-budget-usd', type=float)
|
||||
parser.add_argument('--max-tool-calls', type=int)
|
||||
parser.add_argument('--max-delegated-tasks', type=int)
|
||||
parser.add_argument('--max-model-calls', type=int)
|
||||
parser.add_argument('--max-session-turns', type=int)
|
||||
parser.add_argument('--response-schema-file')
|
||||
parser.add_argument('--response-schema-name')
|
||||
parser.add_argument('--response-schema-strict', action='store_true')
|
||||
@@ -258,6 +265,8 @@ def _build_resumed_agent(args: argparse.Namespace) -> tuple[LocalCodingAgent, St
|
||||
or args.max_budget_usd is not None
|
||||
or args.max_tool_calls is not None
|
||||
or args.max_delegated_tasks is not None
|
||||
or args.max_model_calls is not None
|
||||
or args.max_session_turns is not None
|
||||
):
|
||||
runtime_config = replace(
|
||||
runtime_config,
|
||||
@@ -297,6 +306,16 @@ def _build_resumed_agent(args: argparse.Namespace) -> tuple[LocalCodingAgent, St
|
||||
if args.max_delegated_tasks is not None
|
||||
else runtime_config.budget_config.max_delegated_tasks
|
||||
),
|
||||
max_model_calls=(
|
||||
args.max_model_calls
|
||||
if args.max_model_calls is not None
|
||||
else runtime_config.budget_config.max_model_calls
|
||||
),
|
||||
max_session_turns=(
|
||||
args.max_session_turns
|
||||
if args.max_session_turns is not None
|
||||
else runtime_config.budget_config.max_session_turns
|
||||
),
|
||||
),
|
||||
)
|
||||
output_schema = _load_output_schema_config(args)
|
||||
@@ -339,6 +358,57 @@ def _print_agent_result(result, *, show_transcript: bool) -> None:
|
||||
print(message.get('content', ''))
|
||||
|
||||
|
||||
def _run_agent_chat_loop(
|
||||
agent: LocalCodingAgent,
|
||||
*,
|
||||
initial_prompt: str | None,
|
||||
resume_session_id: str | None,
|
||||
show_transcript: bool,
|
||||
input_func: Callable[[str], str] = input,
|
||||
output_func: Callable[[str], None] = print,
|
||||
result_printer: Callable[..., None] = _print_agent_result,
|
||||
) -> int:
|
||||
active_session_id = resume_session_id
|
||||
first_prompt = initial_prompt
|
||||
|
||||
output_func('# Agent Chat')
|
||||
output_func("Enter a prompt. Use '/exit' or '/quit' to stop.")
|
||||
if active_session_id:
|
||||
output_func(f'resuming_session_id={active_session_id}')
|
||||
|
||||
while True:
|
||||
if first_prompt is not None:
|
||||
prompt = first_prompt
|
||||
first_prompt = None
|
||||
else:
|
||||
try:
|
||||
prompt = input_func('user> ')
|
||||
except EOFError:
|
||||
output_func('chat_ended=eof')
|
||||
return 0
|
||||
except KeyboardInterrupt:
|
||||
output_func('\nchat_ended=interrupt')
|
||||
return 130
|
||||
|
||||
normalized = prompt.strip()
|
||||
if not normalized:
|
||||
continue
|
||||
if normalized in {'/exit', '/quit'}:
|
||||
output_func('chat_ended=user_exit')
|
||||
return 0
|
||||
|
||||
if active_session_id:
|
||||
stored_session = load_agent_session(
|
||||
active_session_id,
|
||||
directory=agent.runtime_config.session_directory,
|
||||
)
|
||||
result = agent.resume(prompt, stored_session)
|
||||
else:
|
||||
result = agent.run(prompt)
|
||||
result_printer(result, show_transcript=show_transcript)
|
||||
active_session_id = result.session_id
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(description='Python porting workspace for the Claude Code rewrite effort')
|
||||
subparsers = parser.add_subparsers(dest='command', required=True)
|
||||
@@ -417,6 +487,13 @@ def build_parser() -> argparse.ArgumentParser:
|
||||
agent_parser.add_argument('--show-transcript', action='store_true')
|
||||
_add_agent_common_args(agent_parser, include_backend=True)
|
||||
|
||||
chat_parser = subparsers.add_parser('agent-chat', help='run an interactive Python local-model chat loop')
|
||||
chat_parser.add_argument('prompt', nargs='?')
|
||||
chat_parser.add_argument('--resume-session-id')
|
||||
chat_parser.add_argument('--max-turns', type=int, default=12)
|
||||
chat_parser.add_argument('--show-transcript', action='store_true')
|
||||
_add_agent_common_args(chat_parser, include_backend=True)
|
||||
|
||||
resume_parser = subparsers.add_parser('agent-resume', help='resume a saved Python local-model agent session')
|
||||
_add_agent_resume_args(resume_parser)
|
||||
|
||||
@@ -563,6 +640,14 @@ def main(argv: list[str] | None = None) -> int:
|
||||
result = agent.run(args.prompt)
|
||||
_print_agent_result(result, show_transcript=args.show_transcript)
|
||||
return 0
|
||||
if args.command == 'agent-chat':
|
||||
agent = _build_agent(args)
|
||||
return _run_agent_chat_loop(
|
||||
agent,
|
||||
initial_prompt=args.prompt,
|
||||
resume_session_id=args.resume_session_id,
|
||||
show_transcript=args.show_transcript,
|
||||
)
|
||||
if args.command == 'agent-resume':
|
||||
agent, stored_session = _build_resumed_agent(args)
|
||||
result = agent.resume(args.prompt, stored_session)
|
||||
|
||||
+219
-6
@@ -44,11 +44,16 @@ class PluginManifest:
|
||||
blocked_tools: tuple[str, ...] = ()
|
||||
before_prompt: str | None = None
|
||||
after_turn: str | None = None
|
||||
on_resume: str | None = None
|
||||
before_persist: str | None = None
|
||||
before_delegate: str | None = None
|
||||
after_delegate: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class PluginRuntime:
|
||||
manifests: tuple[PluginManifest, ...] = field(default_factory=tuple)
|
||||
session_state: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@classmethod
|
||||
def from_workspace(
|
||||
@@ -94,18 +99,83 @@ class PluginRuntime:
|
||||
return tuple(blocks)
|
||||
|
||||
def before_prompt_injections(self) -> tuple[str, ...]:
|
||||
return tuple(
|
||||
if not self.manifests:
|
||||
return ()
|
||||
injections = tuple(
|
||||
manifest.before_prompt
|
||||
for manifest in self.manifests
|
||||
if manifest.before_prompt
|
||||
)
|
||||
if injections:
|
||||
self.session_state['before_prompt_calls'] = int(
|
||||
self.session_state.get('before_prompt_calls', 0)
|
||||
) + 1
|
||||
return injections
|
||||
|
||||
def after_turn_injections(self) -> tuple[str, ...]:
|
||||
return tuple(
|
||||
if not self.manifests:
|
||||
return ()
|
||||
injections = tuple(
|
||||
manifest.after_turn
|
||||
for manifest in self.manifests
|
||||
if manifest.after_turn
|
||||
)
|
||||
if injections:
|
||||
self.session_state['after_turn_calls'] = int(
|
||||
self.session_state.get('after_turn_calls', 0)
|
||||
) + 1
|
||||
return injections
|
||||
|
||||
def on_resume_injections(self) -> tuple[str, ...]:
|
||||
if not self.manifests:
|
||||
return ()
|
||||
injections = tuple(
|
||||
manifest.on_resume
|
||||
for manifest in self.manifests
|
||||
if manifest.on_resume
|
||||
)
|
||||
if injections:
|
||||
self.session_state['resume_calls'] = int(
|
||||
self.session_state.get('resume_calls', 0)
|
||||
) + 1
|
||||
return injections
|
||||
|
||||
def before_persist_injections(self) -> tuple[str, ...]:
|
||||
if not self.manifests:
|
||||
return ()
|
||||
injections = tuple(
|
||||
manifest.before_persist
|
||||
for manifest in self.manifests
|
||||
if manifest.before_persist
|
||||
)
|
||||
if injections:
|
||||
self.session_state['persist_calls'] = int(
|
||||
self.session_state.get('persist_calls', 0)
|
||||
) + 1
|
||||
return injections
|
||||
|
||||
def before_delegate_injections(self) -> tuple[str, ...]:
|
||||
if not self.manifests:
|
||||
return ()
|
||||
injections = tuple(
|
||||
manifest.before_delegate
|
||||
for manifest in self.manifests
|
||||
if manifest.before_delegate
|
||||
)
|
||||
if injections:
|
||||
self.session_state['delegate_calls'] = int(
|
||||
self.session_state.get('delegate_calls', 0)
|
||||
) + 1
|
||||
return injections
|
||||
|
||||
def after_delegate_injections(self) -> tuple[str, ...]:
|
||||
if not self.manifests:
|
||||
return ()
|
||||
return tuple(
|
||||
manifest.after_delegate
|
||||
for manifest in self.manifests
|
||||
if manifest.after_delegate
|
||||
)
|
||||
|
||||
def register_tool_aliases(
|
||||
self,
|
||||
@@ -198,6 +268,111 @@ class PluginRuntime:
|
||||
lines.append(f"- {'; '.join(details)}")
|
||||
if len(self.manifests) > 10:
|
||||
lines.append(f'- ... plus {len(self.manifests) - 10} more plugin manifests')
|
||||
if self.session_state:
|
||||
lines.append(
|
||||
'- runtime_state='
|
||||
+ ', '.join(
|
||||
f'{name}={value}'
|
||||
for name, value in sorted(self.session_state.items())
|
||||
if isinstance(value, (int, float, str, bool))
|
||||
)
|
||||
)
|
||||
return '\n'.join(lines)
|
||||
|
||||
def record_tool_attempt(self, tool_name: str, *, blocked: bool) -> None:
|
||||
attempts = int(self.session_state.get('tool_attempts', 0))
|
||||
self.session_state['tool_attempts'] = attempts + 1
|
||||
if blocked:
|
||||
blocked_count = int(self.session_state.get('blocked_tool_attempts', 0))
|
||||
self.session_state['blocked_tool_attempts'] = blocked_count + 1
|
||||
counts = self.session_state.get('tool_attempt_counts')
|
||||
if not isinstance(counts, dict):
|
||||
counts = {}
|
||||
counts[tool_name] = int(counts.get(tool_name, 0)) + 1
|
||||
self.session_state['tool_attempt_counts'] = counts
|
||||
|
||||
def record_tool_result(
|
||||
self,
|
||||
tool_name: str,
|
||||
*,
|
||||
ok: bool,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
counts = self.session_state.get('tool_result_counts')
|
||||
if not isinstance(counts, dict):
|
||||
counts = {}
|
||||
counts[tool_name] = int(counts.get(tool_name, 0)) + 1
|
||||
self.session_state['tool_result_counts'] = counts
|
||||
key = 'successful_tool_results' if ok else 'failed_tool_results'
|
||||
self.session_state[key] = int(self.session_state.get(key, 0)) + 1
|
||||
if isinstance(metadata, dict) and metadata.get('action') == 'plugin_virtual_tool':
|
||||
self.session_state['virtual_tool_results'] = int(
|
||||
self.session_state.get('virtual_tool_results', 0)
|
||||
) + 1
|
||||
|
||||
def export_session_state(self) -> dict[str, Any]:
|
||||
exported: dict[str, Any] = {}
|
||||
for key, value in self.session_state.items():
|
||||
if isinstance(value, dict):
|
||||
exported[key] = {
|
||||
str(name): count
|
||||
for name, count in value.items()
|
||||
if isinstance(name, str)
|
||||
and isinstance(count, int)
|
||||
and not isinstance(count, bool)
|
||||
}
|
||||
elif isinstance(value, (int, float, str, bool)):
|
||||
exported[key] = value
|
||||
return exported
|
||||
|
||||
def restore_session_state(self, payload: dict[str, Any] | None) -> None:
|
||||
if not isinstance(payload, dict):
|
||||
self.session_state = {}
|
||||
return
|
||||
restored: dict[str, Any] = {}
|
||||
for key, value in payload.items():
|
||||
if isinstance(value, dict):
|
||||
restored[key] = {
|
||||
str(name): int(count)
|
||||
for name, count in value.items()
|
||||
if isinstance(name, str)
|
||||
and isinstance(count, int)
|
||||
and not isinstance(count, bool)
|
||||
}
|
||||
elif isinstance(value, (int, float, str, bool)):
|
||||
restored[key] = value
|
||||
self.session_state = restored
|
||||
|
||||
def runtime_state_reminder(self) -> str | None:
|
||||
if not self.manifests or not self.session_state:
|
||||
return None
|
||||
lines = ['Plugin runtime state:']
|
||||
before_prompt_calls = self.session_state.get('before_prompt_calls')
|
||||
if isinstance(before_prompt_calls, int):
|
||||
lines.append(f'- before_prompt_calls={before_prompt_calls}')
|
||||
after_turn_calls = self.session_state.get('after_turn_calls')
|
||||
if isinstance(after_turn_calls, int):
|
||||
lines.append(f'- after_turn_calls={after_turn_calls}')
|
||||
tool_attempts = self.session_state.get('tool_attempts')
|
||||
if isinstance(tool_attempts, int):
|
||||
lines.append(f'- tool_attempts={tool_attempts}')
|
||||
blocked_attempts = self.session_state.get('blocked_tool_attempts')
|
||||
if isinstance(blocked_attempts, int):
|
||||
lines.append(f'- blocked_tool_attempts={blocked_attempts}')
|
||||
resume_calls = self.session_state.get('resume_calls')
|
||||
if isinstance(resume_calls, int):
|
||||
lines.append(f'- resume_calls={resume_calls}')
|
||||
persist_calls = self.session_state.get('persist_calls')
|
||||
if isinstance(persist_calls, int):
|
||||
lines.append(f'- persist_calls={persist_calls}')
|
||||
delegate_calls = self.session_state.get('delegate_calls')
|
||||
if isinstance(delegate_calls, int):
|
||||
lines.append(f'- delegate_calls={delegate_calls}')
|
||||
virtual_results = self.session_state.get('virtual_tool_results')
|
||||
if isinstance(virtual_results, int):
|
||||
lines.append(f'- virtual_tool_results={virtual_results}')
|
||||
if len(lines) == 1:
|
||||
return None
|
||||
return '\n'.join(lines)
|
||||
|
||||
|
||||
@@ -248,7 +423,15 @@ def _load_manifest(path: Path) -> PluginManifest | None:
|
||||
name = payload.get('name')
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
return None
|
||||
before_prompt, after_turn, hook_names = _parse_hooks(payload.get('hooks'))
|
||||
(
|
||||
before_prompt,
|
||||
after_turn,
|
||||
on_resume,
|
||||
before_persist,
|
||||
before_delegate,
|
||||
after_delegate,
|
||||
hook_names,
|
||||
) = _parse_hooks(payload.get('hooks'))
|
||||
return PluginManifest(
|
||||
name=name.strip(),
|
||||
path=str(path),
|
||||
@@ -266,6 +449,10 @@ def _load_manifest(path: Path) -> PluginManifest | None:
|
||||
),
|
||||
before_prompt=before_prompt,
|
||||
after_turn=after_turn,
|
||||
on_resume=on_resume,
|
||||
before_persist=before_persist,
|
||||
before_delegate=before_delegate,
|
||||
after_delegate=after_delegate,
|
||||
)
|
||||
|
||||
|
||||
@@ -311,10 +498,20 @@ def _extract_tool_aliases(payload: dict[str, Any]) -> tuple[PluginToolAlias, ...
|
||||
return tuple(aliases)
|
||||
|
||||
|
||||
def _parse_hooks(value: Any) -> tuple[str | None, str | None, tuple[str, ...]]:
|
||||
def _parse_hooks(
|
||||
value: Any,
|
||||
) -> tuple[
|
||||
str | None,
|
||||
str | None,
|
||||
str | None,
|
||||
str | None,
|
||||
str | None,
|
||||
str | None,
|
||||
tuple[str, ...],
|
||||
]:
|
||||
if isinstance(value, list):
|
||||
names = tuple(item for item in value if isinstance(item, str) and item.strip())
|
||||
return None, None, names
|
||||
return None, None, None, None, None, None, names
|
||||
if isinstance(value, dict):
|
||||
names = tuple(key for key in value if isinstance(key, str) and key.strip())
|
||||
before_prompt = value.get('beforePrompt')
|
||||
@@ -323,12 +520,28 @@ def _parse_hooks(value: Any) -> tuple[str | None, str | None, tuple[str, ...]]:
|
||||
after_turn = value.get('afterTurn')
|
||||
if after_turn is None:
|
||||
after_turn = value.get('after_turn')
|
||||
on_resume = value.get('onResume')
|
||||
if on_resume is None:
|
||||
on_resume = value.get('on_resume')
|
||||
before_persist = value.get('beforePersist')
|
||||
if before_persist is None:
|
||||
before_persist = value.get('before_persist')
|
||||
before_delegate = value.get('beforeDelegate')
|
||||
if before_delegate is None:
|
||||
before_delegate = value.get('before_delegate')
|
||||
after_delegate = value.get('afterDelegate')
|
||||
if after_delegate is None:
|
||||
after_delegate = value.get('after_delegate')
|
||||
return (
|
||||
_optional_string(before_prompt),
|
||||
_optional_string(after_turn),
|
||||
_optional_string(on_resume),
|
||||
_optional_string(before_persist),
|
||||
_optional_string(before_delegate),
|
||||
_optional_string(after_delegate),
|
||||
names,
|
||||
)
|
||||
return None, None, ()
|
||||
return None, None, None, None, None, None, ()
|
||||
|
||||
|
||||
def _extract_tool_hooks(payload: dict[str, Any]) -> tuple[PluginToolHook, ...]:
|
||||
|
||||
+56
-6
@@ -53,6 +53,7 @@ class QueryEnginePort:
|
||||
transcript_store: TranscriptStore = field(default_factory=TranscriptStore)
|
||||
runtime_agent: LocalCodingAgent | None = None
|
||||
plugin_runtime: PluginRuntime | None = None
|
||||
runtime_cumulative_usage: UsageSummary = field(default_factory=UsageSummary)
|
||||
runtime_event_counts: dict[str, int] = field(default_factory=dict)
|
||||
runtime_message_kind_counts: dict[str, int] = field(default_factory=dict)
|
||||
runtime_mutation_counts: dict[str, int] = field(default_factory=dict)
|
||||
@@ -84,6 +85,7 @@ class QueryEnginePort:
|
||||
session_id=stored.session_id,
|
||||
mutable_messages=list(stored.messages),
|
||||
total_usage=UsageSummary(stored.input_tokens, stored.output_tokens),
|
||||
runtime_cumulative_usage=UsageSummary(stored.input_tokens, stored.output_tokens),
|
||||
transcript_store=transcript,
|
||||
plugin_runtime=PluginRuntime.from_workspace(Path.cwd()),
|
||||
)
|
||||
@@ -115,16 +117,34 @@ class QueryEnginePort:
|
||||
) -> TurnResult:
|
||||
if self.config.use_runtime_agent and self.runtime_agent is not None:
|
||||
result = self._submit_runtime_message(prompt)
|
||||
cumulative_usage = UsageSummary(
|
||||
input_tokens=result.usage.input_tokens,
|
||||
output_tokens=result.usage.output_tokens,
|
||||
)
|
||||
usage = cumulative_usage
|
||||
if self.runtime_cumulative_usage.input_tokens or self.runtime_cumulative_usage.output_tokens:
|
||||
usage = UsageSummary(
|
||||
input_tokens=max(
|
||||
cumulative_usage.input_tokens - self.runtime_cumulative_usage.input_tokens,
|
||||
0,
|
||||
),
|
||||
output_tokens=max(
|
||||
cumulative_usage.output_tokens - self.runtime_cumulative_usage.output_tokens,
|
||||
0,
|
||||
),
|
||||
)
|
||||
else:
|
||||
usage = UsageSummary(
|
||||
input_tokens=cumulative_usage.input_tokens,
|
||||
output_tokens=cumulative_usage.output_tokens,
|
||||
)
|
||||
turn = TurnResult(
|
||||
prompt=prompt,
|
||||
output=result.final_output,
|
||||
matched_commands=matched_commands,
|
||||
matched_tools=matched_tools,
|
||||
permission_denials=denied_tools,
|
||||
usage=UsageSummary(
|
||||
input_tokens=result.usage.input_tokens,
|
||||
output_tokens=result.usage.output_tokens,
|
||||
),
|
||||
usage=usage,
|
||||
stop_reason=result.stop_reason or 'completed',
|
||||
session_id=result.session_id,
|
||||
session_path=result.session_path,
|
||||
@@ -133,7 +153,12 @@ class QueryEnginePort:
|
||||
events=result.events,
|
||||
transcript=result.transcript,
|
||||
)
|
||||
self._record_turn(prompt, turn, denied_tools)
|
||||
self._record_turn(
|
||||
prompt,
|
||||
turn,
|
||||
denied_tools,
|
||||
runtime_cumulative_usage=cumulative_usage,
|
||||
)
|
||||
return turn
|
||||
|
||||
if len(self.mutable_messages) >= self.config.max_turns:
|
||||
@@ -345,6 +370,7 @@ class QueryEnginePort:
|
||||
prompt: str,
|
||||
turn: TurnResult,
|
||||
denied_tools: tuple[PermissionDenial, ...],
|
||||
runtime_cumulative_usage: UsageSummary | None = None,
|
||||
) -> None:
|
||||
self.mutable_messages.append(prompt)
|
||||
self.transcript_store.append(prompt, kind='prompt')
|
||||
@@ -352,7 +378,11 @@ class QueryEnginePort:
|
||||
if self.config.use_runtime_agent:
|
||||
self._record_runtime_turn(turn)
|
||||
self.permission_denials.extend(denied_tools)
|
||||
self.total_usage = turn.usage
|
||||
if runtime_cumulative_usage is not None:
|
||||
self.runtime_cumulative_usage = runtime_cumulative_usage
|
||||
self.total_usage = runtime_cumulative_usage
|
||||
else:
|
||||
self.total_usage = turn.usage
|
||||
self.last_turn = turn
|
||||
if turn.session_id is not None:
|
||||
self.session_id = turn.session_id
|
||||
@@ -537,6 +567,16 @@ class QueryEnginePort:
|
||||
self.runtime_context_reduction.get('preserved_tail_messages', 0)
|
||||
+ preserved_tail_count
|
||||
)
|
||||
max_source_mutation_serial = metadata.get('max_source_mutation_serial')
|
||||
if (
|
||||
isinstance(max_source_mutation_serial, int)
|
||||
and not isinstance(max_source_mutation_serial, bool)
|
||||
):
|
||||
current = self.runtime_context_reduction.get('max_source_mutation_serial', 0)
|
||||
self.runtime_context_reduction['max_source_mutation_serial'] = max(
|
||||
current,
|
||||
max_source_mutation_serial,
|
||||
)
|
||||
compacted_lineage_ids = metadata.get('compacted_lineage_ids')
|
||||
if isinstance(compacted_lineage_ids, list):
|
||||
self.runtime_context_reduction['compacted_lineages'] = (
|
||||
@@ -572,6 +612,16 @@ class QueryEnginePort:
|
||||
if isinstance(revision_count, int) and not isinstance(revision_count, bool):
|
||||
current = self.runtime_lineage_stats.get('max_revision_count', 0)
|
||||
self.runtime_lineage_stats['max_revision_count'] = max(current, revision_count)
|
||||
max_mutation_serial = metadata.get('max_mutation_serial')
|
||||
if (
|
||||
isinstance(max_mutation_serial, int)
|
||||
and not isinstance(max_mutation_serial, bool)
|
||||
):
|
||||
current = self.runtime_lineage_stats.get('max_mutation_serial', 0)
|
||||
self.runtime_lineage_stats['max_mutation_serial'] = max(
|
||||
current,
|
||||
max_mutation_serial,
|
||||
)
|
||||
|
||||
kind = metadata.get('kind')
|
||||
if kind == 'snipped_message':
|
||||
|
||||
@@ -64,6 +64,8 @@ class StoredAgentSession:
|
||||
usage: JSONDict
|
||||
total_cost_usd: float
|
||||
file_history: tuple[JSONDict, ...]
|
||||
budget_state: JSONDict
|
||||
plugin_state: JSONDict
|
||||
scratchpad_directory: str | None = None
|
||||
|
||||
|
||||
@@ -95,6 +97,16 @@ def load_agent_session(session_id: str, directory: Path | None = None) -> Stored
|
||||
file_history=tuple(
|
||||
entry for entry in data.get('file_history', []) if isinstance(entry, dict)
|
||||
),
|
||||
budget_state=(
|
||||
dict(data.get('budget_state', {}))
|
||||
if isinstance(data.get('budget_state'), dict)
|
||||
else {}
|
||||
),
|
||||
plugin_state=(
|
||||
dict(data.get('plugin_state', {}))
|
||||
if isinstance(data.get('plugin_state'), dict)
|
||||
else {}
|
||||
),
|
||||
scratchpad_directory=(
|
||||
str(data['scratchpad_directory'])
|
||||
if isinstance(data.get('scratchpad_directory'), str)
|
||||
@@ -155,6 +167,8 @@ def serialize_runtime_config(runtime_config: AgentRuntimeConfig) -> JSONDict:
|
||||
'max_total_cost_usd': runtime_config.budget_config.max_total_cost_usd,
|
||||
'max_tool_calls': runtime_config.budget_config.max_tool_calls,
|
||||
'max_delegated_tasks': runtime_config.budget_config.max_delegated_tasks,
|
||||
'max_model_calls': runtime_config.budget_config.max_model_calls,
|
||||
'max_session_turns': runtime_config.budget_config.max_session_turns,
|
||||
},
|
||||
'output_schema': (
|
||||
{
|
||||
@@ -205,6 +219,8 @@ def deserialize_runtime_config(payload: JSONDict) -> AgentRuntimeConfig:
|
||||
max_total_cost_usd=_optional_float(budget_payload.get('max_total_cost_usd')),
|
||||
max_tool_calls=_optional_int(budget_payload.get('max_tool_calls')),
|
||||
max_delegated_tasks=_optional_int(budget_payload.get('max_delegated_tasks')),
|
||||
max_model_calls=_optional_int(budget_payload.get('max_model_calls')),
|
||||
max_session_turns=_optional_int(budget_payload.get('max_session_turns')),
|
||||
),
|
||||
output_schema=_deserialize_output_schema(output_schema_payload),
|
||||
session_directory=Path(str(payload.get('session_directory', DEFAULT_AGENT_SESSION_DIR))).resolve(),
|
||||
|
||||
Reference in New Issue
Block a user