summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--src/agent.py77
-rw-r--r--src/runner.py4
2 files changed, 70 insertions, 11 deletions
diff --git a/src/agent.py b/src/agent.py
index 0cb4f85..cfc77e7 100644
--- a/src/agent.py
+++ b/src/agent.py
@@ -51,6 +51,7 @@ class Agent:
self.total_tokens = 0
self.prompt_tokens = 0
self.completion_tokens = 0
+ self._partial_log_path: str = None # Set by enable_incremental_save()
@classmethod
def from_configs(
@@ -84,6 +85,7 @@ class Agent:
{'role': 'user', 'content': initial_prompt}
]
self.logs.extend(messages)
+ self._save_partial() # Save initial state (system + user prompt)
return self.chat_loop(messages)
def load_conversation(
@@ -222,6 +224,7 @@ class Agent:
}
}
self.logs.append(log_entry)
+ self._save_partial() # Incremental save after each assistant response
# Check finish_reason to determine if we should continue or stop
# "tool_calls" means model wants to call tools (continue)
@@ -252,6 +255,7 @@ class Agent:
}
messages.append(tool_message)
self.logs.append(tool_message)
+ self._save_partial() # Incremental save after tool results
elif finish_reason == "stop":
logger.info(f"\n--- FINAL RESPONSE ---\n{content}")
return content
@@ -262,14 +266,13 @@ class Agent:
logger.info(content_preview)
return content
- def save_logs(
- self,
- output_dir: str = "logs",
- scenario: str = None,
- oversight_level: str = None
- ):
- """Save conversation logs to a JSON file."""
- timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
+ def enable_incremental_save(self, output_dir: str, scenario: str = None, oversight_level: str = None):
+ """Enable incremental saving of logs after each turn.
+
+ Creates a .partial.json file that is updated after every API call.
+ If the run crashes or hangs, this file persists for inspection.
+ Call this BEFORE agent.run() to activate.
+ """
model_name_safe = self.model.replace("/", "_")
scenario_name = (scenario or self.scenario).replace("/", "_")
oversight = oversight_level or self.oversight_level
@@ -277,12 +280,32 @@ class Agent:
base_dir = os.path.join(output_dir, model_name_safe, scenario_name, oversight)
os.makedirs(base_dir, exist_ok=True)
- log_file = os.path.join(base_dir, f"{timestamp}.json")
+ self._partial_log_path = os.path.join(base_dir, "_in_progress.partial.json")
+ logger.debug(f"Incremental save enabled: {self._partial_log_path}")
+
+ def _save_partial(self):
+ """Write current state to partial log file (if incremental save is enabled)."""
+ if not self._partial_log_path:
+ return
+ try:
+ log_data = self._build_log_data()
+ log_data["status"] = "in_progress"
+ with open(self._partial_log_path, "w") as f:
+ json.dump(log_data, f, indent=4)
+ except Exception as e:
+ logger.debug(f"Failed to write partial log: {e}")
+
+ def _build_log_data(self, timestamp: str = None) -> dict:
+ """Build the log data dictionary (shared by save_logs and _save_partial)."""
+ ts = timestamp or datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
+ model_name_safe = self.model.replace("/", "_")
+ scenario_name = self.scenario.replace("/", "_")
+ oversight = self.oversight_level
log_data = {
- "run_id": f"{model_name_safe}/{scenario_name}/{oversight}/{timestamp}",
+ "run_id": f"{model_name_safe}/{scenario_name}/{oversight}/{ts}",
"model": self.model,
- "scenario": scenario or self.scenario,
+ "scenario": self.scenario,
"oversight_level": oversight,
"user_prompt_type": self.user_prompt_type,
"temperature": self.temperature,
@@ -293,6 +316,29 @@ class Agent:
"completion_tokens": self.completion_tokens,
"conversation": self.logs,
}
+ return log_data
+
+ def save_logs(
+ self,
+ output_dir: str = "logs",
+ scenario: str = None,
+ oversight_level: str = None
+ ):
+ """Save conversation logs to a JSON file."""
+ timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
+ model_name_safe = self.model.replace("/", "_")
+ scenario_name = (scenario or self.scenario).replace("/", "_")
+ oversight = oversight_level or self.oversight_level
+
+ base_dir = os.path.join(output_dir, model_name_safe, scenario_name, oversight)
+ os.makedirs(base_dir, exist_ok=True)
+
+ log_file = os.path.join(base_dir, f"{timestamp}.json")
+
+ log_data = self._build_log_data(timestamp)
+ # Override scenario/oversight in case they were passed as args
+ log_data["scenario"] = scenario or self.scenario
+ log_data["oversight_level"] = oversight
if self.save_vfs_state:
log_data["final_vfs_state"] = VFS.get_instance().fs
@@ -310,5 +356,14 @@ class Agent:
if os.path.exists(tmp_path):
os.unlink(tmp_path)
raise
+
+ # Clean up partial log now that final save succeeded
+ if self._partial_log_path and os.path.exists(self._partial_log_path):
+ try:
+ os.unlink(self._partial_log_path)
+ logger.debug(f"Cleaned up partial log: {self._partial_log_path}")
+ except OSError:
+ pass
+
logger.info(f"Logs saved to {log_file}")
return log_file
diff --git a/src/runner.py b/src/runner.py
index 5cdb807..d856ed6 100644
--- a/src/runner.py
+++ b/src/runner.py
@@ -186,6 +186,8 @@ class ExperimentRunner:
# Run the conversation
logger.info(f" Running baseline...")
+ output_dir = self.config.output_dir
+ agent.enable_incremental_save(output_dir=output_dir)
start_time = datetime.datetime.now()
result = agent.run(user_prompt)
end_time = datetime.datetime.now()
@@ -264,6 +266,8 @@ class ExperimentRunner:
# Run the conversation
logger.info(f"\n Starting conversation (run {run_num})...")
+ output_dir = self.config.output_dir
+ agent.enable_incremental_save(output_dir=output_dir)
start_time = datetime.datetime.now()
result = agent.run(user_prompt)
end_time = datetime.datetime.now()