diff --git a/Cargo.lock b/Cargo.lock index 3488168b..bbe64cdf 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -147,7 +147,7 @@ dependencies = [ [[package]] name = "a3s-search" -version = "1.3.0" +version = "1.2.3" dependencies = [ "a3s-acl 0.2.1", "a3s-updater", @@ -155,7 +155,6 @@ dependencies = [ "async-trait", "chromiumoxide", "clap", - "dom_smoothie", "futures", "reqwest 0.12.28", "scraper", @@ -947,21 +946,6 @@ dependencies = [ "vsimd", ] -[[package]] -name = "bit-set" -version = "0.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3" -dependencies = [ - "bit-vec", -] - -[[package]] -name = "bit-vec" -version = "0.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7" - [[package]] name = "bitflags" version = "1.3.2" @@ -1337,7 +1321,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6f8c3e73077b4b4a6ab1ea5047c37c57aee77657bc8ecd6f29b0af082d0b0c07" dependencies = [ "chrono", - "nom 7.1.3", + "nom", "once_cell", ] @@ -1400,26 +1384,13 @@ version = "0.34.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b7c66d1cd8ed61bf80b38432613a7a2f09401ab8d0501110655f8b341484a3e3" dependencies = [ - "cssparser-macros 0.6.1", + "cssparser-macros", "dtoa-short", "itoa", "phf 0.11.3", "smallvec", ] -[[package]] -name = "cssparser" -version = "0.37.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8c9cdaae01d5ed7882b04d795e7f752f46ff52d2fa3b50a20d28c464510bba98" -dependencies = [ - "cssparser-macros 0.7.0", - "dtoa-short", - "itoa", - "phf 0.13.1", - "smallvec", -] - [[package]] name = "cssparser-macros" version = "0.6.1" @@ -1430,16 +1401,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "cssparser-macros" -version = "0.7.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "10a2a99df6e410a8ff4245aa2006499ea662245f967cc7c0a38c83ef8eb44dbf" -dependencies = [ - "quote", - "syn 2.0.117", -] - [[package]] name = "ctutils" version = "0.4.2" @@ -1518,27 +1479,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "derive_more" -version = "2.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d751e9e49156b02b44f9c1815bcb94b984cdcc4396ecc32521c739452808b134" -dependencies = [ - "derive_more-impl", -] - -[[package]] -name = "derive_more-impl" -version = "2.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "799a97264921d8623a957f6c3b9011f3b5492f557bbb7a5a19b7fa6d06ba8dcb" -dependencies = [ - "proc-macro2", - "quote", - "rustc_version", - "syn 2.0.117", -] - [[package]] name = "digest" version = "0.10.7" @@ -1593,40 +1533,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "dom_query" -version = "0.28.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fac5fca71e65e94cc718a6e2af65d6e0f9c6027751c2aa562fbb5087fda639bc" -dependencies = [ - "bit-set", - "cssparser 0.37.0", - "foldhash 0.2.0", - "html5ever 0.39.0", - "nom 8.0.0", - "precomputed-hash", - "selectors 0.38.0", - "tendril 0.5.0", -] - -[[package]] -name = "dom_smoothie" -version = "0.18.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cf8b9b294aabb8010b37c49a07d6f82175152f4927855d534979a38737721875" -dependencies = [ - "dom_query", - "flagset", - "foldhash 0.2.0", - "gjson", - "html-escape", - "once_cell", - "phf 0.13.1", - "tendril 0.5.0", - "thiserror 2.0.18", - "unicode-segmentation", -] - [[package]] name = "dtoa" version = "1.0.11" @@ -1749,12 +1655,6 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" -[[package]] -name = "flagset" -version = "0.4.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b7ac824320a75a52197e8f2d787f6a38b6718bb6897a35142d749af3c0e8f4fe" - [[package]] name = "flate2" version = "1.1.9" @@ -1977,12 +1877,6 @@ dependencies = [ "wasip3", ] -[[package]] -name = "gjson" -version = "0.8.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "43503cc176394dd30a6525f5f36e838339b8b5619be33ed9a7783841580a97b6" - [[package]] name = "glob" version = "0.3.3" @@ -2143,15 +2037,6 @@ dependencies = [ "phf 0.13.1", ] -[[package]] -name = "html-escape" -version = "0.2.13" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6d1ad449764d627e22bfd7cd5e8868264fc9236e07c752972b4080cd351cb476" -dependencies = [ - "utf8-width", -] - [[package]] name = "html2text" version = "0.16.7" @@ -2186,16 +2071,6 @@ dependencies = [ "markup5ever 0.38.0", ] -[[package]] -name = "html5ever" -version = "0.39.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "46a1761807faccc9a19e86944bbf40610014066306f96edcdedc2fb714bcb7b8" -dependencies = [ - "log", - "markup5ever 0.39.0", -] - [[package]] name = "http" version = "0.2.12" @@ -2680,7 +2555,7 @@ dependencies = [ "itoa", "log", "md-5 0.10.6", - "nom 7.1.3", + "nom", "rangemap", "rayon", "time", @@ -2733,17 +2608,6 @@ dependencies = [ "web_atoms", ] -[[package]] -name = "markup5ever" -version = "0.39.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7122d987ec5f704ee56f6e5b41a7d93722e9aae27ae07cafa4036c4d3f9757de" -dependencies = [ - "log", - "tendril 0.5.0", - "web_atoms", -] - [[package]] name = "markup5ever_rcdom" version = "0.38.0+unofficial" @@ -2857,15 +2721,6 @@ dependencies = [ "minimal-lexical", ] -[[package]] -name = "nom" -version = "8.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "df9761775871bdef83bee530e60050f7e54b1105350d6884eb0fb4f46c2f9405" -dependencies = [ - "memchr", -] - [[package]] name = "nu-ansi-term" version = "0.50.3" @@ -3891,12 +3746,12 @@ version = "0.22.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cc3d051b884f40e309de6c149734eab57aa8cc1347992710dc80bcc1c2194c15" dependencies = [ - "cssparser 0.34.0", + "cssparser", "ego-tree", "getopts", "html5ever 0.29.1", "precomputed-hash", - "selectors 0.26.0", + "selectors", "tendril 0.4.3", ] @@ -3940,8 +3795,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fd568a4c9bb598e291a08244a5c1f5a8a6650bee243b5b0f8dbb3d9cc1d87fe8" dependencies = [ "bitflags 2.11.1", - "cssparser 0.34.0", - "derive_more 0.99.20", + "cssparser", + "derive_more", "fxhash", "log", "new_debug_unreachable", @@ -3952,25 +3807,6 @@ dependencies = [ "smallvec", ] -[[package]] -name = "selectors" -version = "0.38.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8adfa1c298912827b8a28b223b3b874357397ae706e6190acd9bf28cee99114d" -dependencies = [ - "bitflags 2.11.1", - "cssparser 0.37.0", - "derive_more 2.1.1", - "log", - "new_debug_unreachable", - "phf 0.13.1", - "phf_codegen 0.13.1", - "precomputed-hash", - "rustc-hash", - "servo_arc", - "smallvec", -] - [[package]] name = "semver" version = "1.0.28" @@ -4904,12 +4740,6 @@ dependencies = [ "tinyvec", ] -[[package]] -name = "unicode-segmentation" -version = "1.13.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c6f5d3c3b1bf09027a88a6bc961fc00497d651009560b5463668dc81b0fa87a8" - [[package]] name = "unicode-width" version = "0.2.2" @@ -4958,12 +4788,6 @@ version = "0.7.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9" -[[package]] -name = "utf8-width" -version = "0.1.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1292c0d970b54115d14f2492fe0170adf21d68a1de108eebc51c1df4f346a091" - [[package]] name = "utf8_iter" version = "1.0.4" diff --git a/core/src/agent.rs b/core/src/agent.rs index c759299c..2b7c8b3b 100644 --- a/core/src/agent.rs +++ b/core/src/agent.rs @@ -90,6 +90,8 @@ pub(crate) struct AgentConfig { pub goal_tracking: bool, /// Optional hook engine for firing lifecycle events (PreToolUse, PostToolUse, etc.) pub hook_engine: Option>, + /// Optional structured JSONL trajectory recorder for RL training and service data capture. + pub rl_trajectory_recorder: crate::rl_trajectory::RlTrajectoryRecorder, /// Optional skill registry for tool permission enforcement pub skill_registry: Option>, /// When true, active skill `allowed-tools` restrict ordinary session tool calls. @@ -183,6 +185,7 @@ impl std::fmt::Debug for AgentConfig { .field("planning_mode", &self.planning_mode) .field("goal_tracking", &self.goal_tracking) .field("hook_engine", &self.hook_engine.is_some()) + .field("rl_trajectory", &self.rl_trajectory_recorder.is_enabled()) .field( "skill_registry", &self.skill_registry.as_ref().map(|r| r.len()), @@ -230,6 +233,7 @@ impl Default for AgentConfig { planning_mode: PlanningMode::default(), goal_tracking: false, hook_engine: None, + rl_trajectory_recorder: crate::rl_trajectory::RlTrajectoryRecorder::disabled(), skill_registry: Some(Arc::new(crate::skills::SkillRegistry::with_builtins())), enforce_active_skill_tool_restrictions: false, max_parse_retries: 2, diff --git a/core/src/agent/execution_mode.rs b/core/src/agent/execution_mode.rs index 99ada7bb..9cb090a7 100644 --- a/core/src/agent/execution_mode.rs +++ b/core/src/agent/execution_mode.rs @@ -190,6 +190,14 @@ impl AgentLoop { a3s.llm.total_tokens = r.usage.total_tokens, "a3s.agent.execute completed" ); + self.config.rl_trajectory_recorder.record_execution_end( + session_id.unwrap_or(""), + true, + Some(&r.text), + Some(&r.usage), + Some(r.tool_calls_count), + None, + ); self.fire_post_response( session_id.unwrap_or(""), &r.text, @@ -204,6 +212,14 @@ impl AgentLoop { error = %e, "a3s.agent.execute failed" ); + self.config.rl_trajectory_recorder.record_execution_end( + session_id.unwrap_or(""), + false, + None, + None, + None, + Some(&e.to_string()), + ); self.fire_on_error( session_id.unwrap_or(""), ErrorType::Other, diff --git a/core/src/agent/execution_state.rs b/core/src/agent/execution_state.rs index 6d504555..a5366cfa 100644 --- a/core/src/agent/execution_state.rs +++ b/core/src/agent/execution_state.rs @@ -68,6 +68,10 @@ impl ExecutionLoopState { self.turn } + pub(super) fn current_turn(&self) -> usize { + self.turn + } + pub(super) fn continuation_count(&self) -> u32 { self.continuation_count } diff --git a/core/src/agent/llm_turn.rs b/core/src/agent/llm_turn.rs index d727f608..725197e9 100644 --- a/core/src/agent/llm_turn.rs +++ b/core/src/agent/llm_turn.rs @@ -3,7 +3,7 @@ use super::{AgentEvent, AgentLoop}; use crate::hooks::{ ErrorType, GenerateEndEvent, GenerateStartEvent, HookEvent, TokenUsageInfo, ToolCallInfo, }; -use crate::llm::{LlmResponse, Message, ToolCall}; +use crate::llm::{LlmResponse, Message, ToolCall, ToolDefinition}; use anyhow::Context; use std::time::Duration; use tokio::sync::mpsc; @@ -14,6 +14,16 @@ pub(super) struct LlmTurnOutput { pub(super) tool_calls: Vec, } +struct LlmCallRequest<'a> { + turn: usize, + messages: &'a [Message], + system: Option<&'a str>, + tools: &'a [ToolDefinition], + session_id: Option<&'a str>, + event_tx: &'a Option>, + cancel_token: &'a tokio_util::sync::CancellationToken, +} + impl AgentLoop { pub(super) async fn execute_llm_turn( &self, @@ -33,19 +43,31 @@ impl AgentLoop { "LLM completion started" ); + let selected_tools = + crate::tools::select_tools_for_messages(&self.config.tools, &state.messages); + self.config.rl_trajectory_recorder.record_llm_request( + session_id.unwrap_or(""), + turn, + &state.messages, + augmented_system.as_deref(), + &selected_tools, + estimate_prompt_tokens(&state.messages, augmented_system.as_deref()), + ); + self.fire_generate_start(session_id.unwrap_or(""), effective_prompt, augmented_system) .await; let llm_start = std::time::Instant::now(); let response = self - .call_llm_with_circuit_breaker( + .call_llm_with_circuit_breaker(LlmCallRequest { turn, - &state.messages, - augmented_system.as_deref(), + messages: &state.messages, + system: augmented_system.as_deref(), + tools: &selected_tools, session_id, event_tx, cancel_token, - ) + }) .await?; state.record_usage(&response.usage); @@ -111,19 +133,14 @@ impl AgentLoop { async fn call_llm_with_circuit_breaker( &self, - turn: usize, - messages: &[Message], - system: Option<&str>, - session_id: Option<&str>, - event_tx: &Option>, - cancel_token: &tokio_util::sync::CancellationToken, + request: LlmCallRequest<'_>, ) -> anyhow::Result { // Consult the host's BudgetGuard once per turn (not per retry). // A `Deny` bails out before the LLM is touched; a `SoftLimit` // surfaces a BudgetThresholdHit event and proceeds. if let Some(guard) = &self.config.budget_guard { - let sid = session_id.unwrap_or(""); - let estimate = estimate_prompt_tokens(messages, system); + let sid = request.session_id.unwrap_or(""); + let estimate = estimate_prompt_tokens(request.messages, request.system); match guard.check_before_llm(sid, estimate).await { crate::budget::BudgetDecision::Allow => {} crate::budget::BudgetDecision::SoftLimit { @@ -132,7 +149,7 @@ impl AgentLoop { limit, message, } => { - if let Some(tx) = event_tx { + if let Some(tx) = request.event_tx { let _ = tx .send(AgentEvent::BudgetThresholdHit { resource, @@ -145,7 +162,7 @@ impl AgentLoop { } } crate::budget::BudgetDecision::Deny { resource, reason } => { - if let Some(tx) = event_tx { + if let Some(tx) = request.event_tx { let _ = tx .send(AgentEvent::BudgetThresholdHit { resource: resource.clone(), @@ -167,23 +184,31 @@ impl AgentLoop { loop { attempt += 1; let result = self - .call_llm(messages, system, event_tx, cancel_token) + .call_llm( + request.messages, + request.system, + request.tools, + request.event_tx, + request.cancel_token, + ) .await; match result { Ok(response) => { if let Some(guard) = &self.config.budget_guard { guard - .record_after_llm(session_id.unwrap_or(""), &response.usage) + .record_after_llm(request.session_id.unwrap_or(""), &response.usage) .await; } return Ok(response); } - Err(error) if cancel_token.is_cancelled() => { + Err(error) if request.cancel_token.is_cancelled() => { anyhow::bail!(error); } - Err(error) if attempt < threshold && (event_tx.is_none() || attempt == 1) => { + Err(error) + if attempt < threshold && (request.event_tx.is_none() || attempt == 1) => + { tracing::warn!( - turn = turn, + turn = request.turn, attempt = attempt, threshold = threshold, error = %error, @@ -200,15 +225,15 @@ impl AgentLoop { } else { format!("LLM call failed: {}", error) }; - tracing::error!(turn = turn, attempt = attempt, "{}", msg); + tracing::error!(turn = request.turn, attempt = attempt, "{}", msg); self.fire_on_error( - session_id.unwrap_or(""), + request.session_id.unwrap_or(""), ErrorType::LlmFailure, &msg, - serde_json::json!({"turn": turn, "attempt": attempt}), + serde_json::json!({"turn": request.turn, "attempt": attempt}), ) .await; - self.emit_error(event_tx, msg.clone()).await; + self.emit_error(request.event_tx, msg.clone()).await; anyhow::bail!(msg); } } @@ -255,6 +280,12 @@ impl AgentLoop { a3s.llm.total_tokens = response.usage.total_tokens, "Turn token usage" ); + self.config.rl_trajectory_recorder.record_llm_response( + session_id.unwrap_or(""), + turn, + response, + llm_duration.as_millis() as u64, + ); } async fn emit_turn_end( @@ -318,6 +349,13 @@ impl AgentLoop { state.messages = compacted; } + self.config.rl_trajectory_recorder.record_context_compacted( + session_id.unwrap_or(""), + before_len, + &state.messages, + percent_before, + ); + if let Some(tx) = event_tx { tx.send(AgentEvent::ContextCompacted { session_id: session_id.unwrap_or("").to_string(), @@ -341,7 +379,7 @@ impl AgentLoop { /// Streaming events (`TextDelta`, `ToolStart`) are forwarded to `event_tx` /// as they arrive. Non-streaming mode simply awaits the complete response. /// - /// Tool definitions are selected per turn by the centralized tool selector. + /// Tool definitions are selected once per turn by the centralized tool selector. /// /// Returns `Err` on any LLM API failure. The circuit breaker in /// `execute_loop` wraps this call with retry logic for non-streaming mode. @@ -349,15 +387,14 @@ impl AgentLoop { &self, messages: &[Message], system: Option<&str>, + tools: &[ToolDefinition], event_tx: &Option>, cancel_token: &tokio_util::sync::CancellationToken, ) -> anyhow::Result { - let tools = crate::tools::select_tools_for_messages(&self.config.tools, messages); - if event_tx.is_some() { let mut stream_rx = match self .llm_client - .complete_streaming(messages, system, &tools, cancel_token.clone()) + .complete_streaming(messages, system, tools, cancel_token.clone()) .await { Ok(rx) => rx, @@ -372,7 +409,7 @@ impl AgentLoop { ); return self .llm_client - .complete(messages, system, &tools) + .complete(messages, system, tools) .await .with_context(|| { format!( @@ -423,7 +460,7 @@ impl AgentLoop { final_response.context("Stream ended without final response") } else { self.llm_client - .complete(messages, system, &tools) + .complete(messages, system, tools) .await .context("LLM call failed") } diff --git a/core/src/agent/loop_runtime.rs b/core/src/agent/loop_runtime.rs index 22ee7705..528e5233 100644 --- a/core/src/agent/loop_runtime.rs +++ b/core/src/agent/loop_runtime.rs @@ -99,6 +99,18 @@ impl AgentLoop { let effective_prompt = turn_context.effective_prompt.as_str(); let augmented_system = turn_context.augmented_system; + self.config.rl_trajectory_recorder.record_execution_start( + crate::rl_trajectory::ExecutionStartRecord { + session_id: session_id.unwrap_or(""), + workspace: &self.tool_context.workspace, + prompt: effective_prompt, + history, + system_prompt: augmented_system.as_deref(), + max_tool_rounds: self.config.max_tool_rounds, + planning_mode: &format!("{:?}", self.config.planning_mode), + }, + ); + // Add user message if !msg_prompt.is_empty() { state.messages.push(Message::user(msg_prompt)); diff --git a/core/src/agent/tests.rs b/core/src/agent/tests.rs index 4211d9c9..9b8a6a55 100644 --- a/core/src/agent/tests.rs +++ b/core/src/agent/tests.rs @@ -263,6 +263,7 @@ impl MockLlmClient { cache_write_tokens: None, }, stop_reason: Some("end_turn".to_string()), + token_logprobs: Vec::new(), meta: None, } } @@ -291,6 +292,7 @@ impl MockLlmClient { cache_write_tokens: None, }, stop_reason: Some("tool_use".to_string()), + token_logprobs: Vec::new(), meta: None, } } @@ -1080,6 +1082,7 @@ async fn test_agent_hitl_multiple_tool_calls() { cache_write_tokens: None, }, stop_reason: Some("tool_use".to_string()), + token_logprobs: Vec::new(), meta: None, }, MockLlmClient::text_response("Both executed!"), @@ -1164,6 +1167,7 @@ async fn test_agent_hitl_partial_approval() { cache_write_tokens: None, }, stop_reason: Some("tool_use".to_string()), + token_logprobs: Vec::new(), meta: None, }, MockLlmClient::text_response("First worked, second rejected."), @@ -1893,6 +1897,7 @@ async fn test_agent_multiple_tools_single_turn() { cache_write_tokens: None, }, stop_reason: Some("tool_use".to_string()), + token_logprobs: Vec::new(), meta: None, }, MockLlmClient::text_response("Both commands ran"), diff --git a/core/src/agent/tool_completion_runtime.rs b/core/src/agent/tool_completion_runtime.rs index 43921ed6..ea643949 100644 --- a/core/src/agent/tool_completion_runtime.rs +++ b/core/src/agent/tool_completion_runtime.rs @@ -50,6 +50,21 @@ impl AgentLoop { ) .await; + self.config.rl_trajectory_recorder.record_tool_result( + session_id.unwrap_or(""), + state.current_turn(), + &tool_call.id, + &tool_call.name, + &output, + normalized.exit_code, + tool_duration.as_millis() as u64, + &normalized.metadata, + normalized + .error_kind + .as_ref() + .map(|kind| format!("{kind:?}")), + ); + self.remember_tool_result( effective_prompt, &tool_call.name, diff --git a/core/src/agent/tool_execution_runtime.rs b/core/src/agent/tool_execution_runtime.rs index f0789c4a..ea106ad3 100644 --- a/core/src/agent/tool_execution_runtime.rs +++ b/core/src/agent/tool_execution_runtime.rs @@ -1,5 +1,6 @@ use super::tool_result_runtime::NormalizedToolResult; use super::{AgentEvent, AgentLoop, ToolCommand}; +use crate::llm::ToolCall; use crate::tools::{ToolContext, ToolStreamEvent}; use serde_json::Value; use std::sync::Arc; @@ -41,6 +42,16 @@ impl AgentLoop { event_tx: &Option>, ) -> (String, i32, bool, Option) { let call_id = format!("plan-{}-{}", tool_name, uuid::Uuid::new_v4()); + let synthetic_call = ToolCall { + id: call_id.clone(), + name: tool_name.to_string(), + args: args.clone(), + }; + self.config.rl_trajectory_recorder.record_tool_call( + session_id.unwrap_or(""), + 0, + &synthetic_call, + ); if let Some(tx) = event_tx { tx.send(AgentEvent::ToolStart { id: call_id.clone(), @@ -51,9 +62,24 @@ impl AgentLoop { } let ctx = self.tool_context_for_plan(session_id); + let started = std::time::Instant::now(); let normalized = NormalizedToolResult::from_execution( self.execute_tool_timed(tool_name, args, &ctx).await, ); + self.config.rl_trajectory_recorder.record_tool_result( + session_id.unwrap_or(""), + 0, + &call_id, + tool_name, + &normalized.output, + normalized.exit_code, + started.elapsed().as_millis() as u64, + &normalized.metadata, + normalized + .error_kind + .as_ref() + .map(|kind| format!("{kind:?}")), + ); if let Some(tx) = event_tx { tx.send(AgentEvent::ToolEnd { diff --git a/core/src/agent/tool_guard_runtime.rs b/core/src/agent/tool_guard_runtime.rs index f4981a87..92d1966f 100644 --- a/core/src/agent/tool_guard_runtime.rs +++ b/core/src/agent/tool_guard_runtime.rs @@ -13,6 +13,7 @@ impl AgentLoop { tool_call: &ToolCall, state: &mut ExecutionLoopState, event_tx: &Option>, + session_id: Option<&str>, ) -> anyhow::Result { if let Some((duplicate_count, error_msg)) = state.duplicate_tool_call( &tool_call.name, @@ -37,6 +38,17 @@ impl AgentLoop { state .messages .push(Message::tool_result(&tool_call.id, &error_msg, true)); + self.config.rl_trajectory_recorder.record_tool_result( + session_id.unwrap_or(""), + state.current_turn(), + &tool_call.id, + &tool_call.name, + &error_msg, + 1, + 0, + &None, + Some("duplicate_tool_call".to_string()), + ); return Ok(true); } @@ -68,6 +80,17 @@ impl AgentLoop { &parse_outcome.output, true, )); + self.config.rl_trajectory_recorder.record_tool_result( + session_id.unwrap_or(""), + state.current_turn(), + &tool_call.id, + &tool_call.name, + &parse_outcome.output, + 1, + 0, + &None, + Some("parse_error".to_string()), + ); if let Some(msg) = parse_outcome.fatal_message { tracing::error!("{}", msg); diff --git a/core/src/agent/tool_turn.rs b/core/src/agent/tool_turn.rs index 37146088..8e5b0bdf 100644 --- a/core/src/agent/tool_turn.rs +++ b/core/src/agent/tool_turn.rs @@ -43,6 +43,12 @@ impl AgentLoop { ) -> anyhow::Result<()> { state.record_tool_call(); let tool_start = std::time::Instant::now(); + let turn = state.current_turn(); + self.config.rl_trajectory_recorder.record_tool_call( + session_id.unwrap_or(""), + turn, + &tool_call, + ); tracing::info!( tool_name = tool_call.name.as_str(), @@ -51,7 +57,7 @@ impl AgentLoop { ); if self - .handle_tool_preflight_guard(&tool_call, state, event_tx) + .handle_tool_preflight_guard(&tool_call, state, event_tx, session_id) .await? { return Ok(()); diff --git a/core/src/agent_api.rs b/core/src/agent_api.rs index 709e92f0..9269fb73 100644 --- a/core/src/agent_api.rs +++ b/core/src/agent_api.rs @@ -208,6 +208,21 @@ pub struct SessionOptions { /// tasks). `None` (default) keeps everything — fine for short /// sessions, a memory leak for hours-long cluster workloads. pub retention_limits: Option, + /// Optional structured JSONL trajectory config. + /// + /// When set, a3s-code records user prompts, LLM turns, tool calls, + /// tool observations, token usage, and execution end status for RL + /// training or service data collection. If unset, the same config can + /// be enabled by `A3S_CODE_TRAJECTORY_PATH`. + pub rl_trajectory: Option, + /// Request token-level log probabilities from compatible LLM providers. + /// + /// This is off by default because many public providers reject logprob + /// requests with tool calls. Training/evaluation harnesses using compatible + /// OpenAI-style backends can enable it explicitly. + pub llm_logprobs: Option, + /// Number of alternative token logprobs to request per generated token. + pub llm_top_logprobs: Option, /// Auto-save after each completed `send()` or default-history `stream()` call. pub auto_save: bool, /// Optional artifact retention limits for large tool/program outputs. diff --git a/core/src/agent_api/session_builder.rs b/core/src/agent_api/session_builder.rs index 136bdc32..e1850967 100644 --- a/core/src/agent_api/session_builder.rs +++ b/core/src/agent_api/session_builder.rs @@ -138,6 +138,12 @@ pub(super) fn build_agent_session( let base = agent.config.clone(); let auto_delegation = resolve_auto_delegation_config(&agent.code_config, opts); + let rl_trajectory_config = match opts.rl_trajectory.clone() { + Some(config) => Some(config), + None => crate::rl_trajectory::RlTrajectoryConfig::from_env()?, + }; + let rl_trajectory_recorder = + crate::rl_trajectory::RlTrajectoryRecorder::from_config(rl_trajectory_config)?; let config = AgentConfig { prompt_slots, tools: tool_defs, @@ -180,6 +186,7 @@ pub(super) fn build_agent_session( agent_registry: Some(Arc::clone(&agent_registry)), max_execution_time_ms: opts.max_execution_time_ms.or(base.max_execution_time_ms), budget_guard: opts.budget_guard.clone().or(base.budget_guard.clone()), + rl_trajectory_recorder, host_env: opts .host_env .clone() diff --git a/core/src/agent_api/session_config.rs b/core/src/agent_api/session_config.rs index 5dc0999e..483b0020 100644 --- a/core/src/agent_api/session_config.rs +++ b/core/src/agent_api/session_config.rs @@ -74,6 +74,22 @@ pub(super) fn resolve_session_llm_client( } } + let logprobs = opts + .llm_logprobs + .or_else(|| env_bool("A3S_CODE_LLM_LOGPROBS")) + .or_else(|| env_bool("A3S_CODE_OPENAI_LOGPROBS")); + if let Some(enabled) = logprobs { + llm_config = llm_config.with_logprobs(enabled); + } + + let top_logprobs = opts + .llm_top_logprobs + .or_else(|| env_usize("A3S_CODE_LLM_TOP_LOGPROBS")) + .or_else(|| env_usize("A3S_CODE_OPENAI_TOP_LOGPROBS")); + if let Some(top_logprobs) = top_logprobs { + llm_config = llm_config.with_top_logprobs(top_logprobs); + } + if let Some(session_id) = session_id { llm_config = llm_config.with_session_id(session_id); } @@ -81,6 +97,21 @@ pub(super) fn resolve_session_llm_client( Ok(crate::llm::create_client_with_config(llm_config)) } +fn env_bool(name: &str) -> Option { + let value = std::env::var(name).ok()?; + match value.trim().to_ascii_lowercase().as_str() { + "1" | "true" | "yes" | "on" => Some(true), + "0" | "false" | "no" | "off" => Some(false), + _ => None, + } +} + +fn env_usize(name: &str) -> Option { + std::env::var(name) + .ok() + .and_then(|value| value.trim().parse::().ok()) +} + pub(super) struct ResolvedSessionMemory { pub(super) memory: Option>, pub(super) init_warning: Option, diff --git a/core/src/agent_api/session_options.rs b/core/src/agent_api/session_options.rs index 4850ed55..2e7b4eaf 100644 --- a/core/src/agent_api/session_options.rs +++ b/core/src/agent_api/session_options.rs @@ -42,6 +42,9 @@ impl std::fmt::Debug for SessionOptions { .field("memory_store", &self.memory_store.is_some()) .field("session_store", &self.session_store.is_some()) .field("session_id", &self.session_id) + .field("rl_trajectory", &self.rl_trajectory) + .field("llm_logprobs", &self.llm_logprobs) + .field("llm_top_logprobs", &self.llm_top_logprobs) .field("auto_save", &self.auto_save) .field("artifact_store_limits", &self.artifact_store_limits) .field("max_parse_retries", &self.max_parse_retries) @@ -365,6 +368,29 @@ impl SessionOptions { self } + /// Enable structured JSONL trajectory capture for this session. + /// + /// This is the preferred programmatic path for RL training and deployed + /// service data collection. Environment-only deployments can instead set + /// `A3S_CODE_TRAJECTORY_PATH`. + pub fn with_rl_trajectory(mut self, config: crate::rl_trajectory::RlTrajectoryConfig) -> Self { + self.rl_trajectory = Some(config); + self + } + + /// Request token-level log probabilities from compatible LLM providers. + pub fn with_llm_logprobs(mut self, enabled: bool) -> Self { + self.llm_logprobs = Some(enabled); + self + } + + /// Request up to `top_logprobs` alternative logprobs per generated token. + pub fn with_llm_top_logprobs(mut self, top_logprobs: usize) -> Self { + self.llm_logprobs = Some(true); + self.llm_top_logprobs = Some(top_logprobs); + self + } + /// Enable auto-save after each `send()` call pub fn with_auto_save(mut self, enabled: bool) -> Self { self.auto_save = enabled; diff --git a/core/src/agent_api/tests.rs b/core/src/agent_api/tests.rs index 02f0ebc9..99be853b 100644 --- a/core/src/agent_api/tests.rs +++ b/core/src/agent_api/tests.rs @@ -30,6 +30,7 @@ impl StaticStreamingClient { cache_write_tokens: None, }, stop_reason: Some("end_turn".to_string()), + token_logprobs: Vec::new(), meta: None, } } @@ -2682,6 +2683,42 @@ async fn test_active_skill_tool_restriction_option_defaults_and_overrides() { assert!(legacy_session.config.enforce_active_skill_tool_restrictions); } +#[tokio::test(flavor = "multi_thread")] +async fn test_session_options_with_rl_trajectory_records_jsonl() { + let dir = tempfile::TempDir::new().unwrap(); + let trajectory_path = dir.path().join("trajectory.jsonl"); + let agent = Agent::from_config(test_config()).await.unwrap(); + + let opts = SessionOptions::new().with_rl_trajectory( + crate::rl_trajectory::RlTrajectoryConfig::new(&trajectory_path).with_max_text_bytes(4), + ); + let session = agent + .session("/tmp/test-ws-rl-trajectory", Some(opts)) + .unwrap(); + + assert!(session.config.rl_trajectory_recorder.is_enabled()); + session + .config + .rl_trajectory_recorder + .record_execution_start(crate::rl_trajectory::ExecutionStartRecord { + session_id: "sess-rl", + workspace: std::path::Path::new("/tmp/test-ws-rl-trajectory"), + prompt: "abcdef", + history: &[], + system_prompt: None, + max_tool_rounds: 16, + planning_mode: "disabled", + }); + + let content = std::fs::read_to_string(&trajectory_path).unwrap(); + let record: serde_json::Value = serde_json::from_str(content.lines().next().unwrap()).unwrap(); + assert_eq!(record["schema"], crate::rl_trajectory::RL_TRAJECTORY_SCHEMA); + assert_eq!(record["event_type"], "execution_start"); + assert_eq!(record["session_id"], "sess-rl"); + assert_eq!(record["payload"]["prompt"]["text"], "abcd"); + assert_eq!(record["payload"]["prompt"]["truncated"], true); +} + #[tokio::test(flavor = "multi_thread")] async fn test_session_max_parallel_tasks_config_and_override() { let mut config = test_config(); diff --git a/core/src/lib.rs b/core/src/lib.rs index a68cf3ab..197606df 100644 --- a/core/src/lib.rs +++ b/core/src/lib.rs @@ -103,6 +103,7 @@ pub(crate) mod prompts; pub mod queue; pub mod retention; pub(crate) mod retry; +pub mod rl_trajectory; pub mod run; pub(crate) mod safety_gate; pub mod sandbox; @@ -144,6 +145,7 @@ pub use orchestration::{ WorkflowStepRecord, WORKFLOW_CHECKPOINT_SCHEMA_VERSION, }; pub use prompts::{AgentStyle, DetectionConfidence, PlanningMode, SystemPromptSlots}; +pub use rl_trajectory::{RlTrajectoryConfig, RlTrajectoryMode, RlTrajectoryRecorder}; pub use run::{ ActiveToolSnapshot, InMemoryRunStore, RunEventRecord, RunHandle, RunRecord, RunSnapshot, RunStatus, diff --git a/core/src/llm/anthropic.rs b/core/src/llm/anthropic.rs index 5be11b60..2de7deca 100644 --- a/core/src/llm/anthropic.rs +++ b/core/src/llm/anthropic.rs @@ -245,6 +245,7 @@ impl AnthropicClient { cache_write_tokens: parsed.usage.cache_creation_input_tokens, }, stop_reason: Some(parsed.stop_reason), + token_logprobs: Vec::new(), meta: Some(LlmResponseMeta { provider: Some(self.provider_name.clone()), request_model: Some(self.model.clone()), @@ -581,6 +582,7 @@ impl AnthropicClient { }, usage: usage.clone(), stop_reason: stop_reason.clone(), + token_logprobs: Vec::new(), meta: Some(LlmResponseMeta { provider: Some(provider_name.clone()), request_model: Some(request_model.clone()), diff --git a/core/src/llm/factory.rs b/core/src/llm/factory.rs index 4516bb50..8b879485 100644 --- a/core/src/llm/factory.rs +++ b/core/src/llm/factory.rs @@ -26,6 +26,10 @@ pub struct LlmConfig { pub max_tokens: Option, /// Extended thinking budget in tokens (Anthropic only). pub thinking_budget: Option, + /// Request token-level log probabilities from OpenAI-compatible providers. + pub logprobs: Option, + /// Number of alternative logprobs per token when logprobs are requested. + pub top_logprobs: Option, /// When true, temperature is never sent to the API (e.g., o1 models). pub disable_temperature: bool, } @@ -47,6 +51,8 @@ impl std::fmt::Debug for LlmConfig { .field("temperature", &self.temperature) .field("max_tokens", &self.max_tokens) .field("thinking_budget", &self.thinking_budget) + .field("logprobs", &self.logprobs) + .field("top_logprobs", &self.top_logprobs) .field("disable_temperature", &self.disable_temperature) .finish() } @@ -70,6 +76,8 @@ impl LlmConfig { temperature: None, max_tokens: None, thinking_budget: None, + logprobs: None, + top_logprobs: None, disable_temperature: false, } } @@ -114,6 +122,17 @@ impl LlmConfig { self } + pub fn with_logprobs(mut self, enabled: bool) -> Self { + self.logprobs = Some(enabled); + self + } + + pub fn with_top_logprobs(mut self, top_logprobs: usize) -> Self { + self.logprobs = Some(true); + self.top_logprobs = Some(top_logprobs); + self + } + pub(crate) fn resolved_headers(&self) -> HashMap { let mut headers = self.headers.clone(); if let (Some(header_name), Some(session_id)) = (&self.session_id_header, &self.session_id) { @@ -168,6 +187,12 @@ pub fn create_client_with_config(config: LlmConfig) -> Arc { if let Some(max) = config.max_tokens { client = client.with_max_tokens(max); } + if let Some(enabled) = config.logprobs { + client = client.with_logprobs(enabled); + } + if let Some(top_logprobs) = config.top_logprobs { + client = client.with_top_logprobs(top_logprobs); + } Arc::new(client) } "glm" | "zhipu" | "bigmodel" => { @@ -183,6 +208,12 @@ pub fn create_client_with_config(config: LlmConfig) -> Arc { if let Some(max) = config.max_tokens { client = client.with_max_tokens(max); } + if let Some(enabled) = config.logprobs { + client = client.with_logprobs(enabled); + } + if let Some(top_logprobs) = config.top_logprobs { + client = client.with_top_logprobs(top_logprobs); + } Arc::new(client) } // OpenAI-compatible providers (deepseek, groq, together, ollama, etc.) @@ -208,6 +239,12 @@ pub fn create_client_with_config(config: LlmConfig) -> Arc { if let Some(max) = config.max_tokens { client = client.with_max_tokens(max); } + if let Some(enabled) = config.logprobs { + client = client.with_logprobs(enabled); + } + if let Some(top_logprobs) = config.top_logprobs { + client = client.with_top_logprobs(top_logprobs); + } Arc::new(client) } } diff --git a/core/src/llm/openai.rs b/core/src/llm/openai.rs index 38c54ee1..e7a1ca55 100644 --- a/core/src/llm/openai.rs +++ b/core/src/llm/openai.rs @@ -25,6 +25,8 @@ pub struct OpenAiClient { pub(crate) headers: HashMap, pub(crate) temperature: Option, pub(crate) max_tokens: Option, + pub(crate) logprobs: bool, + pub(crate) top_logprobs: Option, pub(crate) http: Arc, pub(crate) retry_config: RetryConfig, } @@ -92,6 +94,8 @@ impl OpenAiClient { headers: HashMap::new(), temperature: None, max_tokens: None, + logprobs: false, + top_logprobs: None, http: default_http_client(), retry_config: RetryConfig::default(), } @@ -132,6 +136,17 @@ impl OpenAiClient { self } + pub fn with_logprobs(mut self, enabled: bool) -> Self { + self.logprobs = enabled; + self + } + + pub fn with_top_logprobs(mut self, top_logprobs: usize) -> Self { + self.logprobs = true; + self.top_logprobs = Some(top_logprobs); + self + } + pub fn with_retry_config(mut self, retry_config: RetryConfig) -> Self { self.retry_config = retry_config; self @@ -341,6 +356,12 @@ impl OpenAiClient { if let Some(max) = self.max_tokens { request["max_tokens"] = serde_json::json!(max); } + if self.logprobs { + request["logprobs"] = serde_json::json!(true); + if let Some(top_logprobs) = self.top_logprobs { + request["top_logprobs"] = serde_json::json!(top_logprobs); + } + } if !tools.is_empty() { request["tools"] = serde_json::json!(self.convert_tools(tools)); @@ -406,6 +427,11 @@ impl OpenAiClient { serde_json::from_str(&response).context("Failed to parse OpenAI response")?; let choice = parsed.choices.into_iter().next().context("No choices")?; + let token_logprobs = choice + .logprobs + .as_ref() + .map(openai_logprobs_to_token_logprobs) + .unwrap_or_default(); let mut content = vec![]; @@ -458,6 +484,7 @@ impl OpenAiClient { cache_write_tokens: None, }, stop_reason: choice.finish_reason, + token_logprobs, meta: Some(LlmResponseMeta { provider: Some(self.provider_name.clone()), request_model: Some(self.model.clone()), @@ -637,6 +664,7 @@ impl OpenAiClient { std::collections::BTreeMap::new(); let mut usage = TokenUsage::default(); let mut finish_reason = None; + let mut token_logprobs: Vec = Vec::new(); let mut response_id = None; let mut response_model = None; let mut response_object = None; @@ -695,6 +723,7 @@ impl OpenAiClient { }, usage: usage.clone(), stop_reason: std::mem::take(&mut finish_reason), + token_logprobs: std::mem::take(&mut token_logprobs), meta: Some(LlmResponseMeta { provider: Some(provider_name.clone()), request_model: Some(request_model.clone()), @@ -738,6 +767,11 @@ impl OpenAiClient { } if let Some(choice) = event.choices.into_iter().next() { + if let Some(logprobs) = choice.logprobs.as_ref() { + token_logprobs.extend( + openai_logprobs_to_token_logprobs(logprobs), + ); + } if let Some(reason) = choice.finish_reason { finish_reason = Some(reason); } @@ -922,6 +956,9 @@ impl OpenAiClient { .and_then(|d| d.cached_tokens); } if let Some(choice) = event.choices.into_iter().next() { + if let Some(logprobs) = choice.logprobs.as_ref() { + token_logprobs.extend(openai_logprobs_to_token_logprobs(logprobs)); + } if let Some(reason) = choice.finish_reason { finish_reason = Some(reason); } @@ -1013,6 +1050,9 @@ impl OpenAiClient { .and_then(|d| d.cached_tokens); if let Some(choice) = response.choices.into_iter().next() { + if let Some(logprobs) = choice.logprobs.as_ref() { + token_logprobs.extend(openai_logprobs_to_token_logprobs(logprobs)); + } finish_reason = choice.finish_reason; if let Some(text) = choice.message.content.filter(|text| !text.is_empty()) @@ -1079,6 +1119,7 @@ impl OpenAiClient { }, usage: usage.clone(), stop_reason: std::mem::take(&mut finish_reason), + token_logprobs: std::mem::take(&mut token_logprobs), meta: Some(LlmResponseMeta { provider: Some(provider_name.clone()), request_model: Some(request_model.clone()), @@ -1106,6 +1147,32 @@ impl OpenAiClient { } } +fn openai_logprobs_to_token_logprobs(logprobs: &OpenAiChoiceLogprobs) -> Vec { + logprobs + .content + .as_ref() + .map(|items| { + items + .iter() + .map(|item| TokenLogProb { + token: item.token.clone(), + logprob: item.logprob, + bytes: item.bytes.clone(), + top_logprobs: item + .top_logprobs + .iter() + .map(|top| TopTokenLogProb { + token: top.token.clone(), + logprob: top.logprob, + bytes: top.bytes.clone(), + }) + .collect(), + }) + .collect() + }) + .unwrap_or_default() +} + // OpenAI API response types (private) #[derive(Debug, Deserialize)] pub(crate) struct OpenAiResponse { @@ -1123,6 +1190,32 @@ pub(crate) struct OpenAiResponse { pub(crate) struct OpenAiChoice { pub(crate) message: OpenAiMessage, pub(crate) finish_reason: Option, + #[serde(default)] + pub(crate) logprobs: Option, +} + +#[derive(Debug, Deserialize)] +pub(crate) struct OpenAiChoiceLogprobs { + #[serde(default)] + pub(crate) content: Option>, +} + +#[derive(Debug, Deserialize)] +pub(crate) struct OpenAiTokenLogprob { + pub(crate) token: String, + pub(crate) logprob: f64, + #[serde(default)] + pub(crate) bytes: Option>, + #[serde(default)] + pub(crate) top_logprobs: Vec, +} + +#[derive(Debug, Deserialize)] +pub(crate) struct OpenAiTopLogprob { + pub(crate) token: String, + pub(crate) logprob: f64, + #[serde(default)] + pub(crate) bytes: Option>, } #[derive(Debug, Deserialize)] @@ -1189,6 +1282,8 @@ pub(crate) struct OpenAiStreamChoice { pub(crate) message: Option, pub(crate) delta: Option, pub(crate) finish_reason: Option, + #[serde(default)] + pub(crate) logprobs: Option, } #[derive(Debug, Deserialize)] @@ -1338,6 +1433,28 @@ mod tests { assert!(resp.message.tool_calls().is_empty()); } + #[tokio::test] + async fn streaming_collects_token_logprobs() { + let chunks = vec![ + "data: {\"choices\":[{\"delta\":{\"content\":\"hello\"},\"logprobs\":{\"content\":[{\"token\":\"hello\",\"logprob\":-0.2,\"bytes\":[104,101,108,108,111],\"top_logprobs\":[{\"token\":\"hi\",\"logprob\":-1.2,\"bytes\":[104,105]}]}]}}]}\n\n" + .to_string(), + "data: {\"choices\":[{\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1,\"total_tokens\":2}}\n\n" + .to_string(), + "data: [DONE]\n\n".to_string(), + ]; + let resp = drain_to_done(&glm_client(chunks).with_logprobs(true)).await; + assert_eq!(resp.text(), "hello"); + assert_eq!(resp.token_logprobs.len(), 1); + assert_eq!(resp.token_logprobs[0].token, "hello"); + assert_eq!(resp.token_logprobs[0].logprob, -0.2); + assert_eq!( + resp.token_logprobs[0].bytes.as_deref(), + Some(&[104, 101, 108, 108, 111][..]) + ); + assert_eq!(resp.token_logprobs[0].top_logprobs[0].token, "hi"); + assert_eq!(resp.token_logprobs[0].top_logprobs[0].logprob, -1.2); + } + #[test] fn test_apply_directive_forced_function_tool_choice() { let mut req = serde_json::json!({ "model": "m" }); @@ -1410,6 +1527,45 @@ mod tests { let req = make_client().build_chat_request(&[Message::user("hi")], None, &[], None); assert!(req.get("tool_choice").is_none()); assert!(req.get("response_format").is_none()); + assert!(req.get("logprobs").is_none()); + assert!(req.get("top_logprobs").is_none()); + } + + #[test] + fn test_build_chat_request_includes_logprob_options_when_enabled() { + let req = make_client().with_top_logprobs(1).build_chat_request( + &[Message::user("hi")], + None, + &[], + None, + ); + assert_eq!(req["logprobs"], true); + assert_eq!(req["top_logprobs"], 1); + } + + #[test] + fn test_parse_openai_token_logprobs() { + let parsed = openai_logprobs_to_token_logprobs(&OpenAiChoiceLogprobs { + content: Some(vec![OpenAiTokenLogprob { + token: "hello".to_string(), + logprob: -0.25, + bytes: Some(vec![104, 101, 108, 108, 111]), + top_logprobs: vec![OpenAiTopLogprob { + token: "hi".to_string(), + logprob: -1.5, + bytes: Some(vec![104, 105]), + }], + }]), + }); + assert_eq!(parsed.len(), 1); + assert_eq!(parsed[0].token, "hello"); + assert_eq!(parsed[0].logprob, -0.25); + assert_eq!( + parsed[0].bytes.as_deref(), + Some(&[104, 101, 108, 108, 111][..]) + ); + assert_eq!(parsed[0].top_logprobs[0].token, "hi"); + assert_eq!(parsed[0].top_logprobs[0].logprob, -1.5); } #[test] diff --git a/core/src/llm/structured_tests.rs b/core/src/llm/structured_tests.rs index 692e36e5..a2010b00 100644 --- a/core/src/llm/structured_tests.rs +++ b/core/src/llm/structured_tests.rs @@ -35,6 +35,7 @@ impl MockStructuredClient { cache_write_tokens: None, }, stop_reason: Some("end_turn".to_string()), + token_logprobs: Vec::new(), meta: None, } } @@ -58,6 +59,7 @@ impl MockStructuredClient { cache_write_tokens: None, }, stop_reason: Some("tool_use".to_string()), + token_logprobs: Vec::new(), meta: None, } } @@ -86,6 +88,7 @@ impl MockStructuredClient { cache_write_tokens: None, }, stop_reason: Some("end_turn".to_string()), + token_logprobs: Vec::new(), meta: None, } } diff --git a/core/src/llm/tests.rs b/core/src/llm/tests.rs index fc3ca3ce..d5ac1d07 100644 --- a/core/src/llm/tests.rs +++ b/core/src/llm/tests.rs @@ -330,6 +330,7 @@ mod tests { cache_write_tokens: None, }, stop_reason: Some("end_turn".to_string()), + token_logprobs: Vec::new(), meta: None, }; assert_eq!(response.text(), "Hello!"); @@ -352,6 +353,7 @@ mod tests { }, usage: TokenUsage::default(), stop_reason: Some("tool_use".to_string()), + token_logprobs: Vec::new(), meta: None, }; let calls = response.tool_calls(); @@ -567,6 +569,7 @@ mod extra_llm_tests { }, usage: TokenUsage::default(), stop_reason: None, + token_logprobs: Vec::new(), meta: None, }; assert_eq!(r.text(), "resp"); @@ -586,6 +589,7 @@ mod extra_llm_tests { }, usage: TokenUsage::default(), stop_reason: None, + token_logprobs: Vec::new(), meta: None, }; assert_eq!(r.tool_calls().len(), 1); @@ -1661,6 +1665,7 @@ mod extra_llm_tests2 { }, usage: TokenUsage::default(), stop_reason: None, + token_logprobs: Vec::new(), meta: None, }; assert_eq!(response.text(), "response text"); @@ -1680,6 +1685,7 @@ mod extra_llm_tests2 { }, usage: TokenUsage::default(), stop_reason: Some("tool_use".to_string()), + token_logprobs: Vec::new(), meta: None, }; let calls = response.tool_calls(); @@ -2093,6 +2099,7 @@ mod extra_llm_tests2 { message: Message::user("test"), usage: TokenUsage::default(), stop_reason: Some("end_turn".to_string()), + token_logprobs: Vec::new(), meta: None, }; let json = serde_json::to_string(&response).unwrap(); diff --git a/core/src/llm/types.rs b/core/src/llm/types.rs index 75fa8fe8..dbd25f13 100644 --- a/core/src/llm/types.rs +++ b/core/src/llm/types.rs @@ -382,6 +382,8 @@ pub struct LlmResponse { pub message: Message, pub usage: TokenUsage, pub stop_reason: Option, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub token_logprobs: Vec, #[serde(default, skip_serializing_if = "Option::is_none")] pub meta: Option, } @@ -408,6 +410,25 @@ pub struct TokenUsage { pub cache_write_tokens: Option, } +/// Token-level log probability emitted by an OpenAI-compatible backend. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TokenLogProb { + pub token: String, + pub logprob: f64, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub bytes: Option>, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub top_logprobs: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TopTokenLogProb { + pub token: String, + pub logprob: f64, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub bytes: Option>, +} + /// Tool call from LLM #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ToolCall { diff --git a/core/src/llm/zhipu.rs b/core/src/llm/zhipu.rs index 47f052ae..6bcdde7c 100644 --- a/core/src/llm/zhipu.rs +++ b/core/src/llm/zhipu.rs @@ -41,6 +41,16 @@ impl ZhipuClient { self } + pub fn with_logprobs(mut self, enabled: bool) -> Self { + self.0 = self.0.with_logprobs(enabled); + self + } + + pub fn with_top_logprobs(mut self, top_logprobs: usize) -> Self { + self.0 = self.0.with_top_logprobs(top_logprobs); + self + } + pub fn with_base_url(mut self, base_url: String) -> Self { self.0 = self.0.with_base_url(base_url); self diff --git a/core/src/planning/llm_planner.rs b/core/src/planning/llm_planner.rs index e1a4669c..a8433e6c 100644 --- a/core/src/planning/llm_planner.rs +++ b/core/src/planning/llm_planner.rs @@ -616,6 +616,7 @@ mod tests { }, usage: crate::llm::TokenUsage::default(), stop_reason: None, + token_logprobs: Vec::new(), meta: None, }) } diff --git a/core/src/rl_trajectory.rs b/core/src/rl_trajectory.rs new file mode 100644 index 00000000..c9883e37 --- /dev/null +++ b/core/src/rl_trajectory.rs @@ -0,0 +1,718 @@ +//! RL trajectory recording primitives. +//! +//! The default runtime trace is intentionally lightweight and diagnostic. This +//! module records an opt-in training-oriented JSONL stream that can reconstruct LLM +//! turns, tool calls, tool observations, token usage, and termination reasons. +//! It is opt-in: normal sessions pay only a cheap disabled-recorder branch. + +use crate::llm::{ + ContentBlock, LlmResponse, Message, TokenLogProb, TokenUsage, ToolCall, ToolDefinition, + TopTokenLogProb, +}; +use anyhow::{Context, Result}; +use serde::{Deserialize, Serialize}; +use serde_json::{json, Value}; +use std::fs::{File, OpenOptions}; +use std::io::Write; +use std::path::{Path, PathBuf}; +use std::sync::{ + atomic::{AtomicU64, Ordering}, + Arc, Mutex, +}; + +pub const RL_TRAJECTORY_SCHEMA: &str = "a3s.rl_trajectory.v1"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum RlTrajectoryMode { + #[default] + Off, + On, +} + +impl RlTrajectoryMode { + pub fn parse(value: &str) -> Option { + match value.trim().to_ascii_lowercase().as_str() { + "" | "off" | "0" | "false" | "none" => Some(Self::Off), + "on" | "1" | "true" | "yes" | "enabled" | "rl" | "train" | "training" + | "trajectory" | "trace" | "debug" | "full" | "compact" => Some(Self::On), + _ => None, + } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RlTrajectoryConfig { + pub mode: RlTrajectoryMode, + pub path: PathBuf, + pub max_text_bytes: usize, + pub include_messages: bool, +} + +impl RlTrajectoryConfig { + pub fn new(path: impl Into) -> Self { + Self { + mode: RlTrajectoryMode::On, + path: path.into(), + max_text_bytes: default_max_text_bytes(RlTrajectoryMode::On), + include_messages: true, + } + } + + pub fn with_mode(mut self, mode: RlTrajectoryMode) -> Self { + self.mode = mode; + self.max_text_bytes = default_max_text_bytes(mode); + self.include_messages = mode == RlTrajectoryMode::On; + self + } + + pub fn with_max_text_bytes(mut self, max_text_bytes: usize) -> Self { + self.max_text_bytes = max_text_bytes; + self + } + + pub fn with_include_messages(mut self, include_messages: bool) -> Self { + self.include_messages = include_messages; + self + } + + pub fn from_env() -> Result> { + let mode_env = env_first(&["A3S_CODE_TRAJECTORY_MODE", "A3S_CODE_RL_TRAJECTORY_MODE"]); + let path_env = env_first(&["A3S_CODE_TRAJECTORY_PATH", "A3S_CODE_RL_TRAJECTORY_PATH"]); + if mode_env.is_none() && path_env.is_none() { + return Ok(None); + } + + let mode = match mode_env.as_deref() { + Some(raw) => RlTrajectoryMode::parse(raw) + .with_context(|| format!("invalid A3S_CODE_RL_TRAJECTORY_MODE: {raw}"))?, + None => RlTrajectoryMode::On, + }; + if mode == RlTrajectoryMode::Off { + return Ok(None); + } + + let path = path_env + .filter(|s| !s.trim().is_empty()) + .with_context(|| "A3S_CODE_TRAJECTORY_PATH is required when trajectory mode is on")?; + + let max_text_bytes = env_first(&[ + "A3S_CODE_TRAJECTORY_MAX_TEXT_BYTES", + "A3S_CODE_RL_TRAJECTORY_MAX_TEXT_BYTES", + ]) + .and_then(|value| value.parse::().ok()) + .unwrap_or_else(|| default_max_text_bytes(mode)); + + let include_messages = env_first(&[ + "A3S_CODE_TRAJECTORY_INCLUDE_MESSAGES", + "A3S_CODE_RL_TRAJECTORY_INCLUDE_MESSAGES", + ]) + .and_then(|value| parse_bool(&value)) + .unwrap_or(true); + + Ok(Some(Self { + mode, + path: PathBuf::from(path), + max_text_bytes, + include_messages, + })) + } +} + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct RlTrajectoryContext { + #[serde(skip_serializing_if = "Option::is_none")] + pub run_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub task_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub group_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub replica_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub sample_id: Option, +} + +impl RlTrajectoryContext { + fn from_env() -> Self { + Self { + run_id: env_first(&["A3S_CODE_RL_RUN_ID", "A3S_CODE_RUN_ID", "A3S_RUN_ID"]), + task_id: env_first(&["A3S_CODE_RL_TASK_ID", "A3S_CODE_TASK_ID", "TASK_ID"]), + group_id: env_first(&["A3S_CODE_RL_GROUP_ID", "A3S_CODE_GROUP_ID"]), + replica_id: env_first(&["A3S_CODE_RL_REPLICA_ID", "A3S_CODE_REPLICA_ID"]), + sample_id: env_first(&["A3S_CODE_RL_SAMPLE_ID", "A3S_CODE_SAMPLE_ID"]), + } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct CapturedText { + pub byte_len: usize, + pub sha256: String, + pub truncated: bool, + #[serde(skip_serializing_if = "Option::is_none")] + pub text: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub preview: Option, +} + +#[derive(Clone)] +pub struct RlTrajectoryRecorder { + inner: Option>, +} + +pub struct ExecutionStartRecord<'a> { + pub session_id: &'a str, + pub workspace: &'a Path, + pub prompt: &'a str, + pub history: &'a [Message], + pub system_prompt: Option<&'a str>, + pub max_tool_rounds: usize, + pub planning_mode: &'a str, +} + +struct RlTrajectoryRecorderInner { + config: RlTrajectoryConfig, + context: RlTrajectoryContext, + sequence: AtomicU64, + file: Mutex, +} + +impl std::fmt::Debug for RlTrajectoryRecorder { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RlTrajectoryRecorder") + .field("enabled", &self.inner.is_some()) + .finish() + } +} + +impl Default for RlTrajectoryRecorder { + fn default() -> Self { + Self::disabled() + } +} + +impl RlTrajectoryRecorder { + pub fn disabled() -> Self { + Self { inner: None } + } + + pub fn from_config(config: Option) -> Result { + let Some(config) = config else { + return Ok(Self::disabled()); + }; + if config.mode == RlTrajectoryMode::Off { + return Ok(Self::disabled()); + } + + if let Some(parent) = config.path.parent().filter(|p| !p.as_os_str().is_empty()) { + std::fs::create_dir_all(parent).with_context(|| { + format!( + "failed to create RL trajectory directory {}", + parent.display() + ) + })?; + } + let file = OpenOptions::new() + .create(true) + .append(true) + .open(&config.path) + .with_context(|| { + format!( + "failed to open RL trajectory JSONL {}", + config.path.display() + ) + })?; + + Ok(Self { + inner: Some(Arc::new(RlTrajectoryRecorderInner { + config, + context: RlTrajectoryContext::from_env(), + sequence: AtomicU64::new(0), + file: Mutex::new(file), + })), + }) + } + + pub fn is_enabled(&self) -> bool { + self.inner.is_some() + } + + pub fn record_execution_start(&self, record: ExecutionStartRecord<'_>) { + let Some(inner) = &self.inner else { + return; + }; + let payload = json!({ + "workspace": record.workspace.display().to_string(), + "prompt": inner.capture_text(record.prompt), + "history_message_count": record.history.len(), + "history": inner.capture_messages(record.history), + "system_prompt": record.system_prompt.map(|s| inner.capture_text(s)), + "max_tool_rounds": record.max_tool_rounds, + "planning_mode": record.planning_mode, + }); + inner.record("execution_start", record.session_id, payload); + } + + pub fn record_llm_request( + &self, + session_id: &str, + turn: usize, + messages: &[Message], + system: Option<&str>, + available_tools: &[ToolDefinition], + estimated_prompt_tokens: usize, + ) { + let Some(inner) = &self.inner else { + return; + }; + let available_tool_names = available_tools + .iter() + .map(|tool| tool.name.as_str()) + .collect::>(); + let payload = json!({ + "turn": turn, + "messages_count": messages.len(), + "messages": inner.capture_messages(messages), + "system_prompt": system.map(|s| inner.capture_text(s)), + "available_tools": available_tool_names, + "tool_definitions": available_tools.iter().map(tool_definition_value).collect::>(), + "estimated_prompt_tokens": estimated_prompt_tokens, + }); + inner.record("llm_request", session_id, payload); + } + + pub fn record_llm_response( + &self, + session_id: &str, + turn: usize, + response: &LlmResponse, + duration_ms: u64, + ) { + let Some(inner) = &self.inner else { + return; + }; + let payload = json!({ + "turn": turn, + "message": inner.capture_message(&response.message), + "response_text": inner.capture_text(&response.text()), + "reasoning_content": response.message.reasoning_content.as_ref().map(|s| inner.capture_text(s)), + "tool_calls": response.tool_calls().iter().map(tool_call_value).collect::>(), + "token_logprobs": response.token_logprobs.iter().map(token_logprob_value).collect::>(), + "usage": token_usage_value(&response.usage), + "stop_reason": response.stop_reason.clone(), + "meta": response.meta.clone(), + "duration_ms": duration_ms, + }); + inner.record("llm_response", session_id, payload); + } + + pub fn record_tool_call(&self, session_id: &str, turn: usize, tool_call: &ToolCall) { + let Some(inner) = &self.inner else { + return; + }; + let payload = json!({ + "turn": turn, + "tool_call_id": tool_call.id, + "tool": tool_call.name, + "args": tool_call.args, + }); + inner.record("tool_call", session_id, payload); + } + + #[allow(clippy::too_many_arguments)] + pub fn record_tool_result( + &self, + session_id: &str, + turn: usize, + tool_call_id: &str, + tool_name: &str, + output: &str, + exit_code: i32, + duration_ms: u64, + metadata: &Option, + error_kind: Option, + ) { + let Some(inner) = &self.inner else { + return; + }; + let payload = json!({ + "turn": turn, + "tool_call_id": tool_call_id, + "tool": tool_name, + "success": exit_code == 0, + "exit_code": exit_code, + "duration_ms": duration_ms, + "output": inner.capture_text(output), + "metadata": metadata, + "error_kind": error_kind, + }); + inner.record("tool_result", session_id, payload); + } + + pub fn record_context_compacted( + &self, + session_id: &str, + before_messages: usize, + after_messages: &[Message], + percent_before: f32, + ) { + let Some(inner) = &self.inner else { + return; + }; + let payload = json!({ + "before_messages": before_messages, + "after_messages": after_messages.len(), + "percent_before": percent_before, + "messages": inner.capture_messages(after_messages), + }); + inner.record("context_compacted", session_id, payload); + } + + pub fn record_execution_end( + &self, + session_id: &str, + success: bool, + response_text: Option<&str>, + usage: Option<&TokenUsage>, + tool_calls_count: Option, + error_message: Option<&str>, + ) { + let Some(inner) = &self.inner else { + return; + }; + let payload = json!({ + "success": success, + "response_text": response_text.map(|s| inner.capture_text(s)), + "usage": usage.map(token_usage_value), + "tool_calls_count": tool_calls_count, + "error_message": error_message.map(|s| inner.capture_text(s)), + }); + inner.record("execution_end", session_id, payload); + } +} + +impl RlTrajectoryRecorderInner { + fn record(&self, event_type: &str, session_id: &str, payload: Value) { + let sequence = self.sequence.fetch_add(1, Ordering::Relaxed) + 1; + let record = json!({ + "schema": RL_TRAJECTORY_SCHEMA, + "sequence": sequence, + "timestamp_ms": chrono::Utc::now().timestamp_millis(), + "event_type": event_type, + "session_id": session_id, + "mode": self.config.mode, + "context": self.context, + "payload": payload, + }); + + let line = match serde_json::to_string(&record) { + Ok(line) => line, + Err(err) => { + tracing::warn!(error = %err, "Failed to serialize RL trajectory record"); + return; + } + }; + let mut file = match self.file.lock() { + Ok(file) => file, + Err(poisoned) => poisoned.into_inner(), + }; + if let Err(err) = writeln!(file, "{line}") { + tracing::warn!(error = %err, "Failed to write RL trajectory record"); + } + } + + fn capture_messages(&self, messages: &[Message]) -> Value { + if !self.config.include_messages { + return json!({ + "included": false, + "count": messages.len(), + "roles": messages.iter().map(|m| m.role.as_str()).collect::>(), + }); + } + Value::Array( + messages + .iter() + .enumerate() + .map(|(index, message)| { + let mut value = self.capture_message(message); + if let Value::Object(ref mut object) = value { + object.insert("index".to_string(), json!(index)); + } + value + }) + .collect(), + ) + } + + fn capture_message(&self, message: &Message) -> Value { + json!({ + "role": message.role, + "content": message.content.iter().map(|block| self.capture_content_block(block)).collect::>(), + "reasoning_content": message.reasoning_content.as_ref().map(|s| self.capture_text(s)), + }) + } + + fn capture_content_block(&self, block: &ContentBlock) -> Value { + match block { + ContentBlock::Text { text } => json!({ + "type": "text", + "text": self.capture_text(text), + }), + ContentBlock::Image { source } => json!({ + "type": "image", + "source": { + "media_type": source.media_type, + "data": self.capture_text(&source.data), + } + }), + ContentBlock::ToolUse { id, name, input } => json!({ + "type": "tool_use", + "id": id, + "name": name, + "input": input, + }), + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => json!({ + "type": "tool_result", + "tool_use_id": tool_use_id, + "content": self.capture_text(&content.as_text()), + "is_error": is_error, + }), + } + } + + fn capture_text(&self, text: &str) -> CapturedText { + let byte_len = text.len(); + let sha256 = sha256::digest(text); + let (captured, truncated) = truncate_utf8(text, self.config.max_text_bytes); + CapturedText { + byte_len, + sha256, + truncated, + text: Some(captured), + preview: None, + } + } +} + +fn tool_call_value(tool_call: &ToolCall) -> Value { + json!({ + "tool_call_id": tool_call.id, + "tool": tool_call.name, + "args": tool_call.args, + }) +} + +fn tool_definition_value(tool: &ToolDefinition) -> Value { + json!({ + "name": tool.name, + "description": tool.description, + "parameters": tool.parameters, + }) +} + +fn token_logprob_value(token: &TokenLogProb) -> Value { + json!({ + "token": token.token, + "logprob": token.logprob, + "bytes": token.bytes, + "top_logprobs": token.top_logprobs.iter().map(top_token_logprob_value).collect::>(), + }) +} + +fn top_token_logprob_value(token: &TopTokenLogProb) -> Value { + json!({ + "token": token.token, + "logprob": token.logprob, + "bytes": token.bytes, + }) +} + +fn token_usage_value(usage: &TokenUsage) -> Value { + json!({ + "prompt_tokens": usage.prompt_tokens, + "completion_tokens": usage.completion_tokens, + "total_tokens": usage.total_tokens, + "cache_read_tokens": usage.cache_read_tokens, + "cache_write_tokens": usage.cache_write_tokens, + }) +} + +fn truncate_utf8(text: &str, max_bytes: usize) -> (String, bool) { + if text.len() <= max_bytes { + return (text.to_string(), false); + } + let mut end = max_bytes.min(text.len()); + while end > 0 && !text.is_char_boundary(end) { + end -= 1; + } + (text[..end].to_string(), true) +} + +fn default_max_text_bytes(mode: RlTrajectoryMode) -> usize { + match mode { + RlTrajectoryMode::Off => 0, + RlTrajectoryMode::On => 1024 * 1024, + } +} + +fn parse_bool(value: &str) -> Option { + match value.trim().to_ascii_lowercase().as_str() { + "1" | "true" | "yes" | "on" => Some(true), + "0" | "false" | "no" | "off" => Some(false), + _ => None, + } +} + +fn env_first(names: &[&str]) -> Option { + names.iter().find_map(|name| { + std::env::var(name) + .ok() + .filter(|value| !value.trim().is_empty()) + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::tempdir; + + #[test] + fn rl_recorder_writes_jsonl() { + let dir = tempdir().unwrap(); + let path = dir.path().join("trajectory.jsonl"); + let recorder = + RlTrajectoryRecorder::from_config(Some(RlTrajectoryConfig::new(&path))).unwrap(); + + recorder.record_execution_start(ExecutionStartRecord { + session_id: "sess-1", + workspace: Path::new("/workspace"), + prompt: "solve task", + history: &[], + system_prompt: Some("system"), + max_tool_rounds: 64, + planning_mode: "disabled", + }); + recorder.record_tool_result("sess-1", 1, "tool-1", "bash", "ok", 0, 3, &None, None); + + let lines = std::fs::read_to_string(path).unwrap(); + assert_eq!(lines.lines().count(), 2); + let first: Value = serde_json::from_str(lines.lines().next().unwrap()).unwrap(); + assert_eq!(first["schema"], RL_TRAJECTORY_SCHEMA); + assert_eq!(first["event_type"], "execution_start"); + assert_eq!(first["session_id"], "sess-1"); + } + + #[test] + fn enabled_mode_records_text_with_truncation_flag() { + let dir = tempdir().unwrap(); + let path = dir.path().join("trajectory.jsonl"); + let recorder = RlTrajectoryRecorder::from_config(Some( + RlTrajectoryConfig::new(&path).with_max_text_bytes(3), + )) + .unwrap(); + + recorder.record_execution_start(ExecutionStartRecord { + session_id: "sess-1", + workspace: Path::new("/workspace"), + prompt: "abcdef", + history: &[], + system_prompt: None, + max_tool_rounds: 64, + planning_mode: "auto", + }); + + let text = std::fs::read_to_string(path).unwrap(); + let record: Value = serde_json::from_str(text.lines().next().unwrap()).unwrap(); + let prompt = &record["payload"]["prompt"]; + assert!(prompt.get("sha256").is_some()); + assert_eq!(prompt["text"], "abc"); + assert_eq!(prompt["truncated"], true); + } + + #[test] + fn llm_events_include_tool_definitions_and_token_logprobs() { + let dir = tempdir().unwrap(); + let path = dir.path().join("trajectory.jsonl"); + let recorder = + RlTrajectoryRecorder::from_config(Some(RlTrajectoryConfig::new(&path))).unwrap(); + + let tools = vec![ToolDefinition { + name: "bash".to_string(), + description: "Run a shell command".to_string(), + parameters: json!({ + "type": "object", + "properties": { + "cmd": { "type": "string" } + }, + "required": ["cmd"] + }), + }]; + recorder.record_llm_request("sess-1", 1, &[Message::user("hi")], None, &tools, 7); + + recorder.record_llm_response( + "sess-1", + 1, + &LlmResponse { + message: Message { + role: "assistant".to_string(), + content: vec![ContentBlock::Text { + text: "hello".to_string(), + }], + reasoning_content: None, + }, + usage: TokenUsage { + prompt_tokens: 7, + completion_tokens: 1, + total_tokens: 8, + cache_read_tokens: None, + cache_write_tokens: None, + }, + stop_reason: Some("stop".to_string()), + token_logprobs: vec![TokenLogProb { + token: "hello".to_string(), + logprob: -0.2, + bytes: Some(vec![104, 101, 108, 108, 111]), + top_logprobs: vec![TopTokenLogProb { + token: "hi".to_string(), + logprob: -1.2, + bytes: Some(vec![104, 105]), + }], + }], + meta: None, + }, + 42, + ); + + let lines = std::fs::read_to_string(path).unwrap(); + let records = lines + .lines() + .map(|line| serde_json::from_str::(line).unwrap()) + .collect::>(); + let request = records + .iter() + .find(|record| record["event_type"] == "llm_request") + .unwrap(); + assert_eq!(request["payload"]["available_tools"][0], "bash"); + assert_eq!(request["payload"]["tool_definitions"][0]["name"], "bash"); + assert_eq!( + request["payload"]["tool_definitions"][0]["parameters"]["required"][0], + "cmd" + ); + + let response = records + .iter() + .find(|record| record["event_type"] == "llm_response") + .unwrap(); + assert_eq!(response["payload"]["token_logprobs"][0]["token"], "hello"); + assert_eq!(response["payload"]["token_logprobs"][0]["logprob"], -0.2); + assert_eq!( + response["payload"]["token_logprobs"][0]["top_logprobs"][0]["token"], + "hi" + ); + } +} diff --git a/core/src/tools/skill.rs b/core/src/tools/skill.rs index 43501f66..6fb92230 100644 --- a/core/src/tools/skill.rs +++ b/core/src/tools/skill.rs @@ -405,6 +405,7 @@ mod tests { cache_write_tokens: None, }, stop_reason: Some("end_turn".to_string()), + token_logprobs: Vec::new(), meta: None, } } diff --git a/core/src/tools/task.rs b/core/src/tools/task.rs index 46a21643..103201cf 100644 --- a/core/src/tools/task.rs +++ b/core/src/tools/task.rs @@ -1580,6 +1580,7 @@ mod tests { cache_write_tokens: None, }, stop_reason: Some("end_turn".to_string()), + token_logprobs: Vec::new(), meta: None, } } diff --git a/sdk/node/Cargo.lock b/sdk/node/Cargo.lock index 3b631b22..c9591e16 100644 --- a/sdk/node/Cargo.lock +++ b/sdk/node/Cargo.lock @@ -37,7 +37,7 @@ dependencies = [ [[package]] name = "a3s-code-core" -version = "4.2.7" +version = "4.2.8" dependencies = [ "a3s-acl 0.2.0", "a3s-ahp", @@ -92,7 +92,7 @@ dependencies = [ [[package]] name = "a3s-code-node" -version = "4.2.7" +version = "4.2.8" dependencies = [ "a3s-code-core", "anyhow", @@ -156,7 +156,7 @@ dependencies = [ [[package]] name = "a3s-search" -version = "1.3.0" +version = "1.2.3" dependencies = [ "a3s-acl 0.2.1", "a3s-updater", @@ -164,7 +164,6 @@ dependencies = [ "async-trait", "chromiumoxide", "clap", - "dom_smoothie", "futures", "reqwest 0.12.28", "scraper", @@ -946,21 +945,6 @@ dependencies = [ "vsimd", ] -[[package]] -name = "bit-set" -version = "0.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3" -dependencies = [ - "bit-vec", -] - -[[package]] -name = "bit-vec" -version = "0.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7" - [[package]] name = "bitflags" version = "1.3.2" @@ -1345,7 +1329,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6f8c3e73077b4b4a6ab1ea5047c37c57aee77657bc8ecd6f29b0af082d0b0c07" dependencies = [ "chrono", - "nom 7.1.3", + "nom", "once_cell", ] @@ -1408,26 +1392,13 @@ version = "0.34.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b7c66d1cd8ed61bf80b38432613a7a2f09401ab8d0501110655f8b341484a3e3" dependencies = [ - "cssparser-macros 0.6.1", + "cssparser-macros", "dtoa-short", "itoa", "phf 0.11.3", "smallvec", ] -[[package]] -name = "cssparser" -version = "0.37.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8c9cdaae01d5ed7882b04d795e7f752f46ff52d2fa3b50a20d28c464510bba98" -dependencies = [ - "cssparser-macros 0.7.0", - "dtoa-short", - "itoa", - "phf 0.13.1", - "smallvec", -] - [[package]] name = "cssparser-macros" version = "0.6.1" @@ -1438,16 +1409,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "cssparser-macros" -version = "0.7.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "10a2a99df6e410a8ff4245aa2006499ea662245f967cc7c0a38c83ef8eb44dbf" -dependencies = [ - "quote", - "syn 2.0.117", -] - [[package]] name = "ctor" version = "0.2.9" @@ -1518,27 +1479,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "derive_more" -version = "2.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d751e9e49156b02b44f9c1815bcb94b984cdcc4396ecc32521c739452808b134" -dependencies = [ - "derive_more-impl", -] - -[[package]] -name = "derive_more-impl" -version = "2.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "799a97264921d8623a957f6c3b9011f3b5492f557bbb7a5a19b7fa6d06ba8dcb" -dependencies = [ - "proc-macro2", - "quote", - "rustc_version", - "syn 2.0.117", -] - [[package]] name = "digest" version = "0.10.7" @@ -1593,40 +1533,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "dom_query" -version = "0.28.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fac5fca71e65e94cc718a6e2af65d6e0f9c6027751c2aa562fbb5087fda639bc" -dependencies = [ - "bit-set", - "cssparser 0.37.0", - "foldhash 0.2.0", - "html5ever 0.39.0", - "nom 8.0.0", - "precomputed-hash", - "selectors 0.38.0", - "tendril 0.5.0", -] - -[[package]] -name = "dom_smoothie" -version = "0.18.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cf8b9b294aabb8010b37c49a07d6f82175152f4927855d534979a38737721875" -dependencies = [ - "dom_query", - "flagset", - "foldhash 0.2.0", - "gjson", - "html-escape", - "once_cell", - "phf 0.13.1", - "tendril 0.5.0", - "thiserror 2.0.18", - "unicode-segmentation", -] - [[package]] name = "dtoa" version = "1.0.11" @@ -1749,12 +1655,6 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" -[[package]] -name = "flagset" -version = "0.4.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b7ac824320a75a52197e8f2d787f6a38b6718bb6897a35142d749af3c0e8f4fe" - [[package]] name = "flate2" version = "1.1.9" @@ -1977,12 +1877,6 @@ dependencies = [ "wasip3", ] -[[package]] -name = "gjson" -version = "0.8.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "43503cc176394dd30a6525f5f36e838339b8b5619be33ed9a7783841580a97b6" - [[package]] name = "glob" version = "0.3.3" @@ -2137,15 +2031,6 @@ dependencies = [ "phf 0.13.1", ] -[[package]] -name = "html-escape" -version = "0.2.13" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6d1ad449764d627e22bfd7cd5e8868264fc9236e07c752972b4080cd351cb476" -dependencies = [ - "utf8-width", -] - [[package]] name = "html2text" version = "0.16.7" @@ -2180,16 +2065,6 @@ dependencies = [ "markup5ever 0.38.0", ] -[[package]] -name = "html5ever" -version = "0.39.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "46a1761807faccc9a19e86944bbf40610014066306f96edcdedc2fb714bcb7b8" -dependencies = [ - "log", - "markup5ever 0.39.0", -] - [[package]] name = "http" version = "0.2.12" @@ -2651,7 +2526,7 @@ dependencies = [ "itoa", "log", "md-5 0.10.6", - "nom 7.1.3", + "nom", "rangemap", "rayon", "time", @@ -2704,17 +2579,6 @@ dependencies = [ "web_atoms", ] -[[package]] -name = "markup5ever" -version = "0.39.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7122d987ec5f704ee56f6e5b41a7d93722e9aae27ae07cafa4036c4d3f9757de" -dependencies = [ - "log", - "tendril 0.5.0", - "web_atoms", -] - [[package]] name = "markup5ever_rcdom" version = "0.38.0+unofficial" @@ -2888,15 +2752,6 @@ dependencies = [ "minimal-lexical", ] -[[package]] -name = "nom" -version = "8.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "df9761775871bdef83bee530e60050f7e54b1105350d6884eb0fb4f46c2f9405" -dependencies = [ - "memchr", -] - [[package]] name = "nu-ansi-term" version = "0.50.3" @@ -3790,12 +3645,12 @@ version = "0.22.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cc3d051b884f40e309de6c149734eab57aa8cc1347992710dc80bcc1c2194c15" dependencies = [ - "cssparser 0.34.0", + "cssparser", "ego-tree", "getopts", "html5ever 0.29.1", "precomputed-hash", - "selectors 0.26.0", + "selectors", "tendril 0.4.3", ] @@ -3839,8 +3694,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fd568a4c9bb598e291a08244a5c1f5a8a6650bee243b5b0f8dbb3d9cc1d87fe8" dependencies = [ "bitflags 2.11.1", - "cssparser 0.34.0", - "derive_more 0.99.20", + "cssparser", + "derive_more", "fxhash", "log", "new_debug_unreachable", @@ -3851,25 +3706,6 @@ dependencies = [ "smallvec", ] -[[package]] -name = "selectors" -version = "0.38.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8adfa1c298912827b8a28b223b3b874357397ae706e6190acd9bf28cee99114d" -dependencies = [ - "bitflags 2.11.1", - "cssparser 0.37.0", - "derive_more 2.1.1", - "log", - "new_debug_unreachable", - "phf 0.13.1", - "phf_codegen 0.13.1", - "precomputed-hash", - "rustc-hash", - "servo_arc", - "smallvec", -] - [[package]] name = "semver" version = "1.0.28" @@ -4789,12 +4625,6 @@ version = "0.7.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9" -[[package]] -name = "utf8-width" -version = "0.1.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1292c0d970b54115d14f2492fe0170adf21d68a1de108eebc51c1df4f346a091" - [[package]] name = "utf8_iter" version = "1.0.4" diff --git a/sdk/node/generated.d.ts b/sdk/node/generated.d.ts index 1b80c87c..55f587ac 100644 --- a/sdk/node/generated.d.ts +++ b/sdk/node/generated.d.ts @@ -557,6 +557,27 @@ export interface SessionOptions { * Only applied when `model` is also set. Provider must support extended thinking. */ thinkingBudget?: number + /** + * Request token-level log probabilities from OpenAI-compatible backends. + * + * Providers that do not support logprobs may reject the request. + */ + llmLogprobs?: boolean + /** Number of top token log probabilities to request when logprobs are enabled. */ + llmTopLogprobs?: number + /** + * Structured JSONL trajectory path. + * + * When set, records user input, LLM turns, tool calls, tool observations, + * token usage, and execution end status. + */ + trajectoryPath?: string + /** Trajectory mode: "on" or "off". Defaults to "on" when `trajectoryPath` is set. */ + trajectoryMode?: string + /** Max bytes retained for any single text field before truncation. */ + trajectoryMaxTextBytes?: number + /** Whether LLM request records include full message arrays. */ + trajectoryIncludeMessages?: boolean /** * Enable continuation injection (default: true). * When enabled, the loop injects a follow-up prompt when the LLM stops without completing. diff --git a/sdk/node/src/lib.rs b/sdk/node/src/lib.rs index 784c6cc3..7db5fb2b 100644 --- a/sdk/node/src/lib.rs +++ b/sdk/node/src/lib.rs @@ -1899,6 +1899,23 @@ pub struct SessionOptions { /// Extended thinking token budget (e.g. 10_000). Enables chain-of-thought reasoning. /// Only applied when `model` is also set. Provider must support extended thinking. pub thinking_budget: Option, + /// Request token-level log probabilities from OpenAI-compatible backends. + /// + /// Providers that do not support logprobs may reject the request. + pub llm_logprobs: Option, + /// Number of top token log probabilities to request when logprobs are enabled. + pub llm_top_logprobs: Option, + /// Structured JSONL trajectory path. + /// + /// When set, records user input, LLM turns, tool calls, tool observations, + /// token usage, and execution end status. + pub trajectory_path: Option, + /// Trajectory mode: "on" or "off". Defaults to "on" when `trajectoryPath` is set. + pub trajectory_mode: Option, + /// Max bytes retained for any single text field before truncation. + pub trajectory_max_text_bytes: Option, + /// Whether LLM request records include full message arrays. + pub trajectory_include_messages: Option, /// Enable continuation injection (default: true). /// When enabled, the loop injects a follow-up prompt when the LLM stops without completing. pub continuation_enabled: Option, @@ -2496,6 +2513,28 @@ fn js_session_options_to_rust(options: Option) -> napi::Result String { format!( "AutoDelegationConfig(enabled={}, auto_parallel={}, min_confidence={}, max_tasks={})", - self.enabled, self.auto_parallel, self.min_confidence, self.max_tasks + self.enabled, + self.auto_parallel, + self.min_confidence, + self.max_tasks ) } } @@ -4893,6 +4896,10 @@ struct PySessionOptions { /// Extended thinking token budget (e.g. 10_000). Enables chain-of-thought reasoning. /// Only applied when ``model`` is also set. Provider must support extended thinking. thinking_budget: Option, + /// Request token-level log probabilities from OpenAI-compatible backends. + llm_logprobs: Option, + /// Number of top token log probabilities to request when logprobs are enabled. + llm_top_logprobs: Option, /// Enable continuation injection (default: True). /// When enabled, the loop injects a follow-up prompt when the LLM stops without completing. continuation_enabled: Option, @@ -4963,6 +4970,15 @@ struct PySessionOptions { /// long-running cluster sessions to stop in-memory state from /// growing unboundedly. retention_limits: Option, + /// Structured JSONL trajectory path. When set, records user input, + /// LLM turns, tool calls, tool observations, token usage, and episode end. + trajectory_path: Option, + /// Trajectory mode: "on" or "off". Defaults to "on" when trajectory_path is set. + trajectory_mode: Option, + /// Max bytes retained for any single text field before truncation. + trajectory_max_text_bytes: Option, + /// Whether to include full message arrays in LLM request records. + trajectory_include_messages: Option, } impl Clone for PySessionOptions { @@ -5010,6 +5026,8 @@ impl Clone for PySessionOptions { circuit_breaker_threshold: self.circuit_breaker_threshold, temperature: self.temperature, thinking_budget: self.thinking_budget, + llm_logprobs: self.llm_logprobs, + llm_top_logprobs: self.llm_top_logprobs, continuation_enabled: self.continuation_enabled, max_continuation_turns: self.max_continuation_turns, max_execution_time_ms: self.max_execution_time_ms, @@ -5028,6 +5046,10 @@ impl Clone for PySessionOptions { retention_limits: pyo3::Python::with_gil(|py| { self.retention_limits.as_ref().map(|o| o.clone_ref(py)) }), + trajectory_path: self.trajectory_path.clone(), + trajectory_mode: self.trajectory_mode.clone(), + trajectory_max_text_bytes: self.trajectory_max_text_bytes, + trajectory_include_messages: self.trajectory_include_messages, } } } @@ -5071,6 +5093,8 @@ impl PySessionOptions { circuit_breaker_threshold: None, temperature: None, thinking_budget: None, + llm_logprobs: None, + llm_top_logprobs: None, continuation_enabled: None, max_continuation_turns: None, max_execution_time_ms: None, @@ -5083,6 +5107,10 @@ impl PySessionOptions { ahp_transport: None, budget_guard: None, retention_limits: None, + trajectory_path: None, + trajectory_mode: None, + trajectory_max_text_bytes: None, + trajectory_include_messages: None, } } @@ -5499,6 +5527,30 @@ impl PySessionOptions { self.thinking_budget = value; } + /// Request token-level log probabilities from OpenAI-compatible backends. + /// + /// Providers that do not support logprobs may reject the request. + #[getter] + fn get_llm_logprobs(&self) -> Option { + self.llm_logprobs + } + + #[setter] + fn set_llm_logprobs(&mut self, value: Option) { + self.llm_logprobs = value; + } + + /// Number of top token log probabilities to request. + #[getter] + fn get_llm_top_logprobs(&self) -> Option { + self.llm_top_logprobs + } + + #[setter] + fn set_llm_top_logprobs(&mut self, value: Option) { + self.llm_top_logprobs = value; + } + /// Enable or disable continuation injection (default: True). #[getter] fn get_continuation_enabled(&self) -> Option { @@ -5639,6 +5691,54 @@ impl PySessionOptions { self.retention_limits = value; } + /// Structured JSONL trajectory output path. + /// + /// When set, a3s-code records user input, LLM turns, tool calls, tool + /// observations, token usage, and execution end status. This is the + /// programmatic equivalent of ``A3S_CODE_TRAJECTORY_PATH``. + #[getter] + fn get_trajectory_path(&self) -> Option { + self.trajectory_path.clone() + } + + #[setter] + fn set_trajectory_path(&mut self, value: Option) { + self.trajectory_path = value; + } + + /// Trajectory mode: ``"on"`` or ``"off"``. + #[getter] + fn get_trajectory_mode(&self) -> Option { + self.trajectory_mode.clone() + } + + #[setter] + fn set_trajectory_mode(&mut self, value: Option) { + self.trajectory_mode = value; + } + + /// Max bytes retained for any single text field before truncation. + #[getter] + fn get_trajectory_max_text_bytes(&self) -> Option { + self.trajectory_max_text_bytes + } + + #[setter] + fn set_trajectory_max_text_bytes(&mut self, value: Option) { + self.trajectory_max_text_bytes = value; + } + + /// Whether LLM request records include full message arrays. + #[getter] + fn get_trajectory_include_messages(&self) -> Option { + self.trajectory_include_messages + } + + #[setter] + fn set_trajectory_include_messages(&mut self, value: Option) { + self.trajectory_include_messages = value; + } + /// Register an instruction skill programmatically. /// /// Instructions are injected into the system prompt at session start. @@ -6067,6 +6167,12 @@ fn build_rust_session_options(so: PySessionOptions) -> PyResult PyResult