mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-10 03:30:59 +00:00
Add LLM conversation events to PipelineEvent for progress.ndjson observability
Restructures agent events from misaligned EventKind+EventData pair into flat AgentEvent enum, enriches AssistantMessage with model/token/tool_call data, adds 8 new PipelineEvent variants (Prompt, AssistantMessage, ToolCallStarted, ToolCallCompleted, SessionError, ContextWindowWarning, LoopDetected, TurnLimitReached), and forwards agent session events to the pipeline emitter in AgentBackend so they appear in progress.ndjson. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> Entire-Checkpoint: d59f9e69008a
This commit is contained in:
parent
0431fdb242
commit
52d3536a1d
11 changed files with 430 additions and 180 deletions
|
|
@ -1,5 +1,5 @@
|
|||
use crate::{
|
||||
AnthropicProfile, EventData, EventKind, GeminiProfile, LocalExecutionEnvironment, OpenAiProfile,
|
||||
AgentEvent, AnthropicProfile, GeminiProfile, LocalExecutionEnvironment, OpenAiProfile,
|
||||
ProviderProfile, Session, SessionConfig, ToolApprovalFn, Turn,
|
||||
};
|
||||
use clap::{Parser, ValueEnum};
|
||||
|
|
@ -348,8 +348,8 @@ pub async fn run() -> anyhow::Result<()> {
|
|||
tokio::spawn(async move {
|
||||
let s = styles;
|
||||
while let Ok(event) = rx.recv().await {
|
||||
match (&event.kind, &event.data) {
|
||||
(EventKind::ToolCallStart, EventData::ToolCall { tool_name, arguments, .. }) => {
|
||||
match &event.event {
|
||||
AgentEvent::ToolCallStarted { tool_name, arguments, .. } => {
|
||||
eprintln!(
|
||||
" {dim}\u{25cf}{reset} {bold}{cyan}{tool_name}{reset}{dim}({args}){reset}",
|
||||
dim = s.dim,
|
||||
|
|
@ -359,12 +359,9 @@ pub async fn run() -> anyhow::Result<()> {
|
|||
args = format_tool_args(arguments, &cwd_str),
|
||||
);
|
||||
}
|
||||
(
|
||||
EventKind::ToolCallEnd,
|
||||
EventData::ToolCallEnd {
|
||||
tool_name, output, is_error, ..
|
||||
},
|
||||
) if verbose => {
|
||||
AgentEvent::ToolCallCompleted {
|
||||
tool_name, output, is_error, ..
|
||||
} if verbose => {
|
||||
let label = if *is_error { "tool error" } else { "tool result" };
|
||||
eprintln!(
|
||||
" {}[{label}] {tool_name}:{}\n{}",
|
||||
|
|
@ -374,7 +371,7 @@ pub async fn run() -> anyhow::Result<()> {
|
|||
.unwrap_or_else(|_| output.to_string()),
|
||||
);
|
||||
}
|
||||
(EventKind::Error, EventData::Error { error }) => {
|
||||
AgentEvent::Error { error } => {
|
||||
eprintln!(
|
||||
" {red}\u{2717} {error}{reset}",
|
||||
red = s.red,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use crate::types::{EventData, EventKind, SessionEvent};
|
||||
use crate::types::{AgentEvent, SessionEvent};
|
||||
use std::time::SystemTime;
|
||||
use tokio::sync::broadcast;
|
||||
|
||||
|
|
@ -14,15 +14,14 @@ impl EventEmitter {
|
|||
Self { sender }
|
||||
}
|
||||
|
||||
pub fn emit(&self, kind: EventKind, session_id: String, data: EventData) {
|
||||
let event = SessionEvent {
|
||||
kind,
|
||||
pub fn emit(&self, session_id: String, event: AgentEvent) {
|
||||
let wrapped = SessionEvent {
|
||||
event,
|
||||
timestamp: SystemTime::now(),
|
||||
session_id,
|
||||
data,
|
||||
};
|
||||
// Ignore send error (no receivers)
|
||||
let _ = self.sender.send(event);
|
||||
let _ = self.sender.send(wrapped);
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
|
|
@ -46,12 +45,11 @@ mod tests {
|
|||
let emitter = EventEmitter::new();
|
||||
let mut receiver = emitter.subscribe();
|
||||
|
||||
emitter.emit(EventKind::SessionStart, "sess-1".into(), EventData::Empty);
|
||||
emitter.emit("sess-1".into(), AgentEvent::SessionStarted);
|
||||
|
||||
let event = receiver.recv().await.unwrap();
|
||||
assert_eq!(event.kind, EventKind::SessionStart);
|
||||
assert!(matches!(event.event, AgentEvent::SessionStarted));
|
||||
assert_eq!(event.session_id, "sess-1");
|
||||
assert!(matches!(event.data, EventData::Empty));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
@ -60,17 +58,15 @@ mod tests {
|
|||
let mut receiver = emitter.subscribe();
|
||||
|
||||
emitter.emit(
|
||||
EventKind::Error,
|
||||
"sess-2".into(),
|
||||
EventData::Error {
|
||||
AgentEvent::Error {
|
||||
error: "something went wrong".into(),
|
||||
},
|
||||
);
|
||||
|
||||
let event = receiver.recv().await.unwrap();
|
||||
assert_eq!(event.kind, EventKind::Error);
|
||||
assert!(
|
||||
matches!(&event.data, EventData::Error { error } if error == "something went wrong")
|
||||
matches!(&event.event, AgentEvent::Error { error } if error == "something went wrong")
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -80,12 +76,12 @@ mod tests {
|
|||
let mut rx1 = emitter.subscribe();
|
||||
let mut rx2 = emitter.subscribe();
|
||||
|
||||
emitter.emit(EventKind::SessionEnd, "sess-3".into(), EventData::Empty);
|
||||
emitter.emit("sess-3".into(), AgentEvent::SessionEnded);
|
||||
|
||||
let e1 = rx1.recv().await.unwrap();
|
||||
let e2 = rx2.recv().await.unwrap();
|
||||
assert_eq!(e1.kind, EventKind::SessionEnd);
|
||||
assert_eq!(e2.kind, EventKind::SessionEnd);
|
||||
assert!(matches!(e1.event, AgentEvent::SessionEnded));
|
||||
assert!(matches!(e2.event, AgentEvent::SessionEnded));
|
||||
assert_eq!(e1.session_id, "sess-3");
|
||||
assert_eq!(e2.session_id, "sess-3");
|
||||
}
|
||||
|
|
@ -94,9 +90,8 @@ mod tests {
|
|||
fn emit_without_subscribers_does_not_panic() {
|
||||
let emitter = EventEmitter::new();
|
||||
emitter.emit(
|
||||
EventKind::Error,
|
||||
"sess-4".into(),
|
||||
EventData::Error {
|
||||
AgentEvent::Error {
|
||||
error: "test".into(),
|
||||
},
|
||||
);
|
||||
|
|
|
|||
|
|
@ -39,7 +39,7 @@ pub use tools::{
|
|||
make_shell_tool_with_config, make_write_file_tool,
|
||||
};
|
||||
pub use truncation::{truncate_lines, truncate_output, truncate_tool_output, TruncationMode};
|
||||
pub use types::{EventData, EventKind, SessionEvent, SessionState, Turn};
|
||||
pub use types::{AgentEvent, SessionEvent, SessionState, Turn};
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) mod test_support;
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ use crate::project_docs::discover_project_docs;
|
|||
use crate::provider_profile::ProviderProfile;
|
||||
use crate::tool_registry::ToolRegistry;
|
||||
use crate::truncation::truncate_tool_output;
|
||||
use crate::types::{EventData, EventKind, SessionState, Turn};
|
||||
use crate::types::{AgentEvent, SessionState, Turn};
|
||||
use std::collections::VecDeque;
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::SystemTime;
|
||||
|
|
@ -65,7 +65,7 @@ impl Session {
|
|||
/// Call before `process_input`.
|
||||
pub async fn initialize(&mut self) {
|
||||
self.event_emitter
|
||||
.emit(EventKind::SessionStart, self.id.clone(), EventData::Empty);
|
||||
.emit(self.id.clone(), AgentEvent::SessionStarted);
|
||||
|
||||
let doc_root = self
|
||||
.config
|
||||
|
|
@ -176,7 +176,7 @@ impl Session {
|
|||
if self.state != SessionState::Closed {
|
||||
self.state = SessionState::Closed;
|
||||
self.event_emitter
|
||||
.emit(EventKind::SessionEnd, self.id.clone(), EventData::Empty);
|
||||
.emit(self.id.clone(), AgentEvent::SessionEnded);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -228,7 +228,7 @@ impl Session {
|
|||
timestamp: SystemTime::now(),
|
||||
});
|
||||
self.event_emitter
|
||||
.emit(EventKind::UserInput, self.id.clone(), EventData::Empty);
|
||||
.emit(self.id.clone(), AgentEvent::UserInput);
|
||||
|
||||
// Drain steering queue before first LLM call
|
||||
self.drain_steering();
|
||||
|
|
@ -247,14 +247,14 @@ impl Session {
|
|||
// Check max_tool_rounds_per_input
|
||||
if round_count >= self.config.max_tool_rounds_per_input {
|
||||
self.event_emitter
|
||||
.emit(EventKind::TurnLimit, self.id.clone(), EventData::Empty);
|
||||
.emit(self.id.clone(), AgentEvent::TurnLimitReached);
|
||||
break;
|
||||
}
|
||||
|
||||
// Check max_turns
|
||||
if self.config.max_turns > 0 && self.history.turns().len() >= self.config.max_turns {
|
||||
self.event_emitter
|
||||
.emit(EventKind::TurnLimit, self.id.clone(), EventData::Empty);
|
||||
.emit(self.id.clone(), AgentEvent::TurnLimitReached);
|
||||
break;
|
||||
}
|
||||
|
||||
|
|
@ -268,20 +268,16 @@ impl Session {
|
|||
let request = self.build_request(&system_prompt);
|
||||
|
||||
// Emit AssistantTextStart before LLM call
|
||||
self.event_emitter.emit(
|
||||
EventKind::AssistantTextStart,
|
||||
self.id.clone(),
|
||||
EventData::Empty,
|
||||
);
|
||||
self.event_emitter
|
||||
.emit(self.id.clone(), AgentEvent::AssistantTextStart);
|
||||
|
||||
// Call LLM (streaming)
|
||||
let mut event_stream = match self.llm_client.stream(&request).await {
|
||||
Ok(stream) => stream,
|
||||
Err(err) => {
|
||||
self.event_emitter.emit(
|
||||
EventKind::Error,
|
||||
self.id.clone(),
|
||||
EventData::Error {
|
||||
AgentEvent::Error {
|
||||
error: err.to_string(),
|
||||
},
|
||||
);
|
||||
|
|
@ -299,9 +295,8 @@ impl Session {
|
|||
Ok(event) => {
|
||||
if let StreamEvent::TextDelta { ref delta, .. } = event {
|
||||
self.event_emitter.emit(
|
||||
EventKind::AssistantTextDelta,
|
||||
self.id.clone(),
|
||||
EventData::TextDelta {
|
||||
AgentEvent::TextDelta {
|
||||
delta: delta.clone(),
|
||||
},
|
||||
);
|
||||
|
|
@ -310,9 +305,8 @@ impl Session {
|
|||
}
|
||||
Err(err) => {
|
||||
self.event_emitter.emit(
|
||||
EventKind::Error,
|
||||
self.id.clone(),
|
||||
EventData::Error {
|
||||
AgentEvent::Error {
|
||||
error: err.to_string(),
|
||||
},
|
||||
);
|
||||
|
|
@ -370,11 +364,16 @@ impl Session {
|
|||
timestamp: SystemTime::now(),
|
||||
});
|
||||
|
||||
// Emit AssistantTextEnd
|
||||
// Emit AssistantMessage with enriched data from the response
|
||||
self.event_emitter.emit(
|
||||
EventKind::AssistantTextEnd,
|
||||
self.id.clone(),
|
||||
EventData::Empty,
|
||||
AgentEvent::AssistantMessage {
|
||||
text: text.clone(),
|
||||
model: response.model.clone(),
|
||||
input_tokens: response.usage.input_tokens,
|
||||
output_tokens: response.usage.output_tokens,
|
||||
tool_call_count: tool_calls.len(),
|
||||
},
|
||||
);
|
||||
|
||||
// Check context window usage
|
||||
|
|
@ -417,11 +416,8 @@ impl Session {
|
|||
content: "WARNING: Loop detected. You appear to be repeating the same tool calls. Please try a different approach or ask for clarification.".to_string(),
|
||||
timestamp: SystemTime::now(),
|
||||
});
|
||||
self.event_emitter.emit(
|
||||
EventKind::LoopDetection,
|
||||
self.id.clone(),
|
||||
EventData::Empty,
|
||||
);
|
||||
self.event_emitter
|
||||
.emit(self.id.clone(), AgentEvent::LoopDetected);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -440,11 +436,8 @@ impl Session {
|
|||
content: msg,
|
||||
timestamp: SystemTime::now(),
|
||||
});
|
||||
self.event_emitter.emit(
|
||||
EventKind::SteeringInjected,
|
||||
self.id.clone(),
|
||||
EventData::Empty,
|
||||
);
|
||||
self.event_emitter
|
||||
.emit(self.id.clone(), AgentEvent::SteeringInjected);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -506,9 +499,8 @@ impl Session {
|
|||
}
|
||||
|
||||
self.event_emitter.emit(
|
||||
EventKind::ToolCallStart,
|
||||
self.id.clone(),
|
||||
EventData::ToolCall {
|
||||
AgentEvent::ToolCallStarted {
|
||||
tool_name: tc.name.clone(),
|
||||
tool_call_id: tc.id.clone(),
|
||||
arguments: tc.arguments.clone(),
|
||||
|
|
@ -527,17 +519,15 @@ impl Session {
|
|||
.await;
|
||||
|
||||
self.event_emitter.emit(
|
||||
EventKind::ToolCallOutputDelta,
|
||||
self.id.clone(),
|
||||
EventData::TextDelta {
|
||||
AgentEvent::ToolCallOutputDelta {
|
||||
delta: result.content.to_string(),
|
||||
},
|
||||
);
|
||||
|
||||
self.event_emitter.emit(
|
||||
EventKind::ToolCallEnd,
|
||||
self.id.clone(),
|
||||
EventData::ToolCallEnd {
|
||||
AgentEvent::ToolCallCompleted {
|
||||
tool_name: tc.name.clone(),
|
||||
tool_call_id: tc.id.clone(),
|
||||
output: result.content.clone(),
|
||||
|
|
@ -574,9 +564,8 @@ impl Session {
|
|||
let tc = tc.clone();
|
||||
async move {
|
||||
emitter.emit(
|
||||
EventKind::ToolCallStart,
|
||||
session_id.clone(),
|
||||
EventData::ToolCall {
|
||||
AgentEvent::ToolCallStarted {
|
||||
tool_name: tc.name.clone(),
|
||||
tool_call_id: tc.id.clone(),
|
||||
arguments: tc.arguments.clone(),
|
||||
|
|
@ -595,17 +584,15 @@ impl Session {
|
|||
.await;
|
||||
|
||||
emitter.emit(
|
||||
EventKind::ToolCallOutputDelta,
|
||||
session_id.clone(),
|
||||
EventData::TextDelta {
|
||||
AgentEvent::ToolCallOutputDelta {
|
||||
delta: result.content.to_string(),
|
||||
},
|
||||
);
|
||||
|
||||
emitter.emit(
|
||||
EventKind::ToolCallEnd,
|
||||
session_id,
|
||||
EventData::ToolCallEnd {
|
||||
AgentEvent::ToolCallCompleted {
|
||||
tool_name: tc.name.clone(),
|
||||
tool_call_id: tc.id.clone(),
|
||||
output: result.content.clone(),
|
||||
|
|
@ -663,9 +650,8 @@ impl Session {
|
|||
|
||||
if estimated_tokens > threshold {
|
||||
self.event_emitter.emit(
|
||||
EventKind::ContextWindowWarning,
|
||||
self.id.clone(),
|
||||
EventData::ContextWarning {
|
||||
AgentEvent::ContextWindowWarning {
|
||||
estimated_tokens,
|
||||
context_window_size: context_window,
|
||||
usage_percent: estimated_tokens * 100 / context_window,
|
||||
|
|
@ -959,13 +945,13 @@ mod tests {
|
|||
// Collect events
|
||||
let mut events = Vec::new();
|
||||
while let Ok(event) = rx.try_recv() {
|
||||
events.push(event.kind.clone());
|
||||
events.push(event);
|
||||
}
|
||||
|
||||
assert!(events.contains(&EventKind::SessionStart));
|
||||
assert!(events.contains(&EventKind::UserInput));
|
||||
assert!(events.contains(&EventKind::AssistantTextEnd));
|
||||
assert!(events.contains(&EventKind::SessionEnd));
|
||||
assert!(events.iter().any(|e| matches!(e.event, AgentEvent::SessionStarted)));
|
||||
assert!(events.iter().any(|e| matches!(e.event, AgentEvent::UserInput)));
|
||||
assert!(events.iter().any(|e| matches!(e.event, AgentEvent::AssistantMessage { .. })));
|
||||
assert!(events.iter().any(|e| matches!(e.event, AgentEvent::SessionEnded)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
@ -985,17 +971,17 @@ mod tests {
|
|||
|
||||
let mut tool_end_events = Vec::new();
|
||||
while let Ok(event) = rx.try_recv() {
|
||||
if event.kind == EventKind::ToolCallEnd {
|
||||
if matches!(event.event, AgentEvent::ToolCallCompleted { .. }) {
|
||||
tool_end_events.push(event);
|
||||
}
|
||||
}
|
||||
|
||||
assert_eq!(tool_end_events.len(), 1);
|
||||
match &tool_end_events[0].data {
|
||||
EventData::ToolCallEnd { output, .. } => {
|
||||
match &tool_end_events[0].event {
|
||||
AgentEvent::ToolCallCompleted { output, .. } => {
|
||||
assert_eq!(output, &serde_json::json!("echo: hello world"));
|
||||
}
|
||||
_ => panic!("Expected ToolCallEnd event data"),
|
||||
_ => panic!("Expected ToolCallCompleted event"),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -1073,10 +1059,10 @@ mod tests {
|
|||
|
||||
session.process_input("Keep echoing").await.unwrap();
|
||||
|
||||
// Check for LoopDetection event
|
||||
// Check for LoopDetected event
|
||||
let mut found_loop_detection = false;
|
||||
while let Ok(event) = rx.try_recv() {
|
||||
if event.kind == EventKind::LoopDetection {
|
||||
if matches!(event.event, AgentEvent::LoopDetected) {
|
||||
found_loop_detection = true;
|
||||
}
|
||||
}
|
||||
|
|
@ -1238,14 +1224,14 @@ mod tests {
|
|||
let result = session.process_input("Hello").await;
|
||||
assert!(matches!(result, Err(AgentError::SessionClosed)));
|
||||
|
||||
// No SessionStart event should have been emitted
|
||||
// No SessionStarted event should have been emitted
|
||||
let mut events = Vec::new();
|
||||
while let Ok(event) = rx.try_recv() {
|
||||
events.push(event.kind.clone());
|
||||
events.push(event);
|
||||
}
|
||||
assert!(
|
||||
!events.contains(&EventKind::SessionStart),
|
||||
"SessionStart should not be emitted for a closed session"
|
||||
!events.iter().any(|e| matches!(e.event, AgentEvent::SessionStarted)),
|
||||
"SessionStarted should not be emitted for a closed session"
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -1289,13 +1275,13 @@ mod tests {
|
|||
panic!("Expected ToolResults turn at index 2");
|
||||
}
|
||||
|
||||
// Verify ToolCallStart and ToolCallEnd events for all 3 calls
|
||||
// Verify ToolCallStarted and ToolCallCompleted events for all 3 calls
|
||||
let mut start_count = 0;
|
||||
let mut end_count = 0;
|
||||
while let Ok(event) = rx.try_recv() {
|
||||
match event.kind {
|
||||
EventKind::ToolCallStart => start_count += 1,
|
||||
EventKind::ToolCallEnd => end_count += 1,
|
||||
match &event.event {
|
||||
AgentEvent::ToolCallStarted { .. } => start_count += 1,
|
||||
AgentEvent::ToolCallCompleted { .. } => end_count += 1,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
|
@ -1327,17 +1313,13 @@ mod tests {
|
|||
|
||||
let mut found_warning = false;
|
||||
while let Ok(event) = rx.try_recv() {
|
||||
if event.kind == EventKind::ContextWindowWarning {
|
||||
if let AgentEvent::ContextWindowWarning {
|
||||
context_window_size,
|
||||
..
|
||||
} = &event.event
|
||||
{
|
||||
found_warning = true;
|
||||
match &event.data {
|
||||
EventData::ContextWarning {
|
||||
context_window_size,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(*context_window_size, 100);
|
||||
}
|
||||
_ => panic!("Expected ContextWarning event data"),
|
||||
}
|
||||
assert_eq!(*context_window_size, 100);
|
||||
}
|
||||
}
|
||||
assert!(found_warning);
|
||||
|
|
@ -1380,7 +1362,7 @@ mod tests {
|
|||
|
||||
let mut found_warning = false;
|
||||
while let Ok(event) = rx.try_recv() {
|
||||
if event.kind == EventKind::ContextWindowWarning {
|
||||
if matches!(event.event, AgentEvent::ContextWindowWarning { .. }) {
|
||||
found_warning = true;
|
||||
}
|
||||
}
|
||||
|
|
@ -1486,14 +1468,14 @@ mod tests {
|
|||
let mut session_start_count = 0;
|
||||
let mut session_end_count = 0;
|
||||
while let Ok(event) = rx.try_recv() {
|
||||
if event.kind == EventKind::SessionStart {
|
||||
if matches!(event.event, AgentEvent::SessionStarted) {
|
||||
session_start_count += 1;
|
||||
}
|
||||
if event.kind == EventKind::SessionEnd {
|
||||
if matches!(event.event, AgentEvent::SessionEnded) {
|
||||
session_end_count += 1;
|
||||
}
|
||||
}
|
||||
// SESSION_START is emitted once during initialize(), SESSION_END once during close()
|
||||
// SessionStarted is emitted once during initialize(), SessionEnded once during close()
|
||||
assert_eq!(session_start_count, 1);
|
||||
assert_eq!(session_end_count, 1);
|
||||
}
|
||||
|
|
@ -1681,17 +1663,17 @@ mod tests {
|
|||
|
||||
let mut tool_end_events = Vec::new();
|
||||
while let Ok(event) = rx.try_recv() {
|
||||
if event.kind == EventKind::ToolCallEnd {
|
||||
if matches!(event.event, AgentEvent::ToolCallCompleted { .. }) {
|
||||
tool_end_events.push(event);
|
||||
}
|
||||
}
|
||||
|
||||
assert_eq!(tool_end_events.len(), 1);
|
||||
match &tool_end_events[0].data {
|
||||
EventData::ToolCallEnd { is_error, .. } => {
|
||||
assert!(is_error, "ToolCallEnd event should have is_error: true");
|
||||
match &tool_end_events[0].event {
|
||||
AgentEvent::ToolCallCompleted { is_error, .. } => {
|
||||
assert!(is_error, "ToolCallCompleted event should have is_error: true");
|
||||
}
|
||||
_ => panic!("Expected ToolCallEnd event data"),
|
||||
_ => panic!("Expected ToolCallCompleted event"),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -1704,10 +1686,8 @@ mod tests {
|
|||
|
||||
let mut deltas = Vec::new();
|
||||
while let Ok(event) = rx.try_recv() {
|
||||
if event.kind == EventKind::AssistantTextDelta {
|
||||
if let EventData::TextDelta { delta } = &event.data {
|
||||
deltas.push(delta.clone());
|
||||
}
|
||||
if let AgentEvent::TextDelta { delta } = &event.event {
|
||||
deltas.push(delta.clone());
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -43,33 +43,31 @@ pub enum SessionState {
|
|||
Closed,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum EventKind {
|
||||
SessionStart,
|
||||
SessionEnd,
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum AgentEvent {
|
||||
SessionStarted,
|
||||
SessionEnded,
|
||||
UserInput,
|
||||
AssistantTextStart,
|
||||
AssistantTextDelta,
|
||||
AssistantTextEnd,
|
||||
ToolCallStart,
|
||||
ToolCallOutputDelta,
|
||||
ToolCallEnd,
|
||||
SteeringInjected,
|
||||
TurnLimit,
|
||||
LoopDetection,
|
||||
ContextWindowWarning,
|
||||
Error,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum EventData {
|
||||
Empty,
|
||||
ToolCall {
|
||||
AssistantMessage {
|
||||
text: String,
|
||||
model: String,
|
||||
input_tokens: i64,
|
||||
output_tokens: i64,
|
||||
tool_call_count: usize,
|
||||
},
|
||||
TextDelta {
|
||||
delta: String,
|
||||
},
|
||||
ToolCallStarted {
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
arguments: serde_json::Value,
|
||||
},
|
||||
ToolCallEnd {
|
||||
ToolCallOutputDelta {
|
||||
delta: String,
|
||||
},
|
||||
ToolCallCompleted {
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
output: serde_json::Value,
|
||||
|
|
@ -78,22 +76,21 @@ pub enum EventData {
|
|||
Error {
|
||||
error: String,
|
||||
},
|
||||
TextDelta {
|
||||
delta: String,
|
||||
},
|
||||
ContextWarning {
|
||||
ContextWindowWarning {
|
||||
estimated_tokens: usize,
|
||||
context_window_size: usize,
|
||||
usage_percent: usize,
|
||||
},
|
||||
LoopDetected,
|
||||
TurnLimitReached,
|
||||
SteeringInjected,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SessionEvent {
|
||||
pub kind: EventKind,
|
||||
pub event: AgentEvent,
|
||||
pub timestamp: SystemTime,
|
||||
pub session_id: String,
|
||||
pub data: EventData,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
@ -103,13 +100,23 @@ mod tests {
|
|||
#[test]
|
||||
fn session_event_construction() {
|
||||
let event = SessionEvent {
|
||||
kind: EventKind::SessionStart,
|
||||
event: AgentEvent::SessionStarted,
|
||||
timestamp: SystemTime::now(),
|
||||
session_id: "sess_1".into(),
|
||||
data: EventData::Empty,
|
||||
};
|
||||
assert_eq!(event.kind, EventKind::SessionStart);
|
||||
assert!(matches!(event.event, AgentEvent::SessionStarted));
|
||||
assert_eq!(event.session_id, "sess_1");
|
||||
assert!(matches!(event.data, EventData::Empty));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn agent_event_assistant_message() {
|
||||
let event = AgentEvent::AssistantMessage {
|
||||
text: "Hello".into(),
|
||||
model: "test-model".into(),
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
tool_call_count: 2,
|
||||
};
|
||||
assert!(matches!(event, AgentEvent::AssistantMessage { tool_call_count: 2, .. }));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ use std::sync::Arc;
|
|||
use async_trait::async_trait;
|
||||
|
||||
use agent::{
|
||||
AnthropicProfile, DockerConfig, DockerExecutionEnvironment, EventData, EventKind,
|
||||
AgentEvent, AnthropicProfile, DockerConfig, DockerExecutionEnvironment,
|
||||
ExecutionEnvironment, GeminiProfile, LocalExecutionEnvironment, OpenAiProfile, ProviderProfile,
|
||||
Session, SessionConfig, Turn,
|
||||
};
|
||||
|
|
@ -62,6 +62,7 @@ impl CodergenBackend for AgentBackend {
|
|||
prompt: &str,
|
||||
_context: &Context,
|
||||
_thread_id: Option<&str>,
|
||||
emitter: &Arc<crate::event::EventEmitter>,
|
||||
) -> Result<CodergenResult, AttractorError> {
|
||||
let client = Client::from_env()
|
||||
.await
|
||||
|
|
@ -90,23 +91,99 @@ impl CodergenBackend for AgentBackend {
|
|||
|
||||
let mut session = Session::new(client, profile, exec_env, config);
|
||||
|
||||
// Subscribe to session events for real-time tool status on stderr.
|
||||
// Subscribe to session events: forward to pipeline emitter and optionally print to stderr.
|
||||
let verbose = self.verbose;
|
||||
if verbose >= 1 {
|
||||
let node_id = node.id.clone();
|
||||
let styles = self.styles;
|
||||
let mut rx = session.subscribe();
|
||||
tokio::spawn(async move {
|
||||
while let Ok(event) = rx.recv().await {
|
||||
match (&event.kind, &event.data) {
|
||||
(
|
||||
EventKind::ToolCallStart,
|
||||
EventData::ToolCall {
|
||||
tool_name,
|
||||
arguments,
|
||||
..
|
||||
},
|
||||
) => {
|
||||
let node_id = node.id.clone();
|
||||
let styles = self.styles;
|
||||
let pipeline_emitter = Arc::clone(emitter);
|
||||
let mut rx = session.subscribe();
|
||||
tokio::spawn(async move {
|
||||
use crate::event::PipelineEvent;
|
||||
while let Ok(event) = rx.recv().await {
|
||||
// Forward agent events to pipeline events
|
||||
match &event.event {
|
||||
AgentEvent::AssistantMessage {
|
||||
text,
|
||||
model,
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
tool_call_count,
|
||||
} => {
|
||||
pipeline_emitter.emit(&PipelineEvent::AssistantMessage {
|
||||
stage: node_id.clone(),
|
||||
text: text.clone(),
|
||||
model: model.clone(),
|
||||
input_tokens: *input_tokens,
|
||||
output_tokens: *output_tokens,
|
||||
tool_call_count: *tool_call_count,
|
||||
});
|
||||
}
|
||||
AgentEvent::ToolCallStarted {
|
||||
tool_name,
|
||||
tool_call_id,
|
||||
arguments,
|
||||
} => {
|
||||
pipeline_emitter.emit(&PipelineEvent::ToolCallStarted {
|
||||
stage: node_id.clone(),
|
||||
tool_name: tool_name.clone(),
|
||||
tool_call_id: tool_call_id.clone(),
|
||||
arguments: arguments.clone(),
|
||||
});
|
||||
}
|
||||
AgentEvent::ToolCallCompleted {
|
||||
tool_name,
|
||||
tool_call_id,
|
||||
output,
|
||||
is_error,
|
||||
} => {
|
||||
pipeline_emitter.emit(&PipelineEvent::ToolCallCompleted {
|
||||
stage: node_id.clone(),
|
||||
tool_name: tool_name.clone(),
|
||||
tool_call_id: tool_call_id.clone(),
|
||||
output: output.clone(),
|
||||
is_error: *is_error,
|
||||
});
|
||||
}
|
||||
AgentEvent::Error { error } => {
|
||||
pipeline_emitter.emit(&PipelineEvent::SessionError {
|
||||
stage: node_id.clone(),
|
||||
error: error.clone(),
|
||||
});
|
||||
}
|
||||
AgentEvent::ContextWindowWarning {
|
||||
estimated_tokens,
|
||||
context_window_size,
|
||||
usage_percent,
|
||||
} => {
|
||||
pipeline_emitter.emit(&PipelineEvent::ContextWindowWarning {
|
||||
stage: node_id.clone(),
|
||||
estimated_tokens: *estimated_tokens,
|
||||
context_window_size: *context_window_size,
|
||||
usage_percent: *usage_percent,
|
||||
});
|
||||
}
|
||||
AgentEvent::LoopDetected => {
|
||||
pipeline_emitter.emit(&PipelineEvent::LoopDetected {
|
||||
stage: node_id.clone(),
|
||||
});
|
||||
}
|
||||
AgentEvent::TurnLimitReached => {
|
||||
pipeline_emitter.emit(&PipelineEvent::TurnLimitReached {
|
||||
stage: node_id.clone(),
|
||||
});
|
||||
}
|
||||
// Streaming events and session lifecycle not forwarded
|
||||
_ => {}
|
||||
}
|
||||
|
||||
// Verbose stderr printing (gated on verbosity)
|
||||
if verbose >= 1 {
|
||||
match &event.event {
|
||||
AgentEvent::ToolCallStarted {
|
||||
tool_name,
|
||||
arguments,
|
||||
..
|
||||
} => {
|
||||
eprintln!(
|
||||
"{dim}[{node_id}]{reset} {dim}\u{25cf}{reset} {bold}{cyan}{tool_name}{reset}{dim}({args}){reset}",
|
||||
dim = styles.dim,
|
||||
|
|
@ -116,15 +193,12 @@ impl CodergenBackend for AgentBackend {
|
|||
args = format_tool_args(arguments),
|
||||
);
|
||||
}
|
||||
(
|
||||
EventKind::ToolCallEnd,
|
||||
EventData::ToolCallEnd {
|
||||
tool_name,
|
||||
output,
|
||||
is_error,
|
||||
..
|
||||
},
|
||||
) if verbose >= 2 => {
|
||||
AgentEvent::ToolCallCompleted {
|
||||
tool_name,
|
||||
output,
|
||||
is_error,
|
||||
..
|
||||
} if verbose >= 2 => {
|
||||
let label = if *is_error { "error" } else { "result" };
|
||||
eprintln!(
|
||||
"{dim}[{node_id}] [{label}] {tool_name}:{reset}\n{}",
|
||||
|
|
@ -134,7 +208,7 @@ impl CodergenBackend for AgentBackend {
|
|||
reset = styles.reset,
|
||||
);
|
||||
}
|
||||
(EventKind::Error, EventData::Error { error }) => {
|
||||
AgentEvent::Error { error } => {
|
||||
eprintln!(
|
||||
"{dim}[{node_id}]{reset} {red}\u{2717} {error}{reset}",
|
||||
dim = styles.dim,
|
||||
|
|
@ -145,8 +219,14 @@ impl CodergenBackend for AgentBackend {
|
|||
_ => {}
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// Emit Prompt event before processing
|
||||
emitter.emit(&crate::event::PipelineEvent::Prompt {
|
||||
stage: node.id.clone(),
|
||||
text: prompt.to_string(),
|
||||
});
|
||||
|
||||
session.initialize().await;
|
||||
session.process_input(prompt).await.map_err(|e| {
|
||||
|
|
|
|||
|
|
@ -269,6 +269,53 @@ pub fn format_event_summary(event: &PipelineEvent, styles: &Styles) -> String {
|
|||
PipelineEvent::CheckpointSaved { node_id } => {
|
||||
format!("[CHECKPOINT_SAVED] node={node_id}")
|
||||
}
|
||||
PipelineEvent::Prompt { stage, text } => {
|
||||
let truncated = if text.len() > 80 { &text[..80] } else { text };
|
||||
format!("[PROMPT] stage={stage} text=\"{truncated}\"")
|
||||
}
|
||||
PipelineEvent::AssistantMessage {
|
||||
stage,
|
||||
model,
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
tool_call_count,
|
||||
..
|
||||
} => {
|
||||
let total = input_tokens + output_tokens;
|
||||
let tokens_str = format_tokens_human(total);
|
||||
format!("[ASSISTANT_MESSAGE] stage={stage} model={model} tokens={tokens_str} tool_calls={tool_call_count}")
|
||||
}
|
||||
PipelineEvent::ToolCallStarted {
|
||||
stage,
|
||||
tool_name,
|
||||
..
|
||||
} => {
|
||||
format!("[TOOL_CALL_STARTED] stage={stage} tool={tool_name}")
|
||||
}
|
||||
PipelineEvent::ToolCallCompleted {
|
||||
stage,
|
||||
tool_name,
|
||||
is_error,
|
||||
..
|
||||
} => {
|
||||
format!("[TOOL_CALL_COMPLETED] stage={stage} tool={tool_name} is_error={is_error}")
|
||||
}
|
||||
PipelineEvent::SessionError { stage, error } => {
|
||||
format!("[SESSION_ERROR] stage={stage} error=\"{error}\"")
|
||||
}
|
||||
PipelineEvent::ContextWindowWarning {
|
||||
stage,
|
||||
usage_percent,
|
||||
..
|
||||
} => {
|
||||
format!("[CONTEXT_WINDOW_WARNING] stage={stage} usage={usage_percent}%")
|
||||
}
|
||||
PipelineEvent::LoopDetected { stage } => {
|
||||
format!("[LOOP_DETECTED] stage={stage}")
|
||||
}
|
||||
PipelineEvent::TurnLimitReached { stage } => {
|
||||
format!("[TURN_LIMIT_REACHED] stage={stage}")
|
||||
}
|
||||
};
|
||||
format!("{dim}{body}{reset}", dim = styles.dim, reset = styles.reset)
|
||||
}
|
||||
|
|
@ -389,6 +436,63 @@ pub fn format_event_detail(event: &PipelineEvent, styles: &Styles) -> String {
|
|||
"{d}── CHECKPOINT_SAVED ─────────────────────────{r}\n {d}node_id:{r} {node_id}\n"
|
||||
)
|
||||
}
|
||||
PipelineEvent::Prompt { stage, text } => {
|
||||
format!("{d}── PROMPT ───────────────────────────────────{r}\n {d}stage:{r} {stage}\n {d}text:{r}\n{text}\n")
|
||||
}
|
||||
PipelineEvent::AssistantMessage {
|
||||
stage,
|
||||
text,
|
||||
model,
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
tool_call_count,
|
||||
} => {
|
||||
let total = input_tokens + output_tokens;
|
||||
let truncated = if text.len() > 200 { &text[..200] } else { text.as_str() };
|
||||
format!("{d}── ASSISTANT_MESSAGE ────────────────────────{r}\n {d}stage:{r} {stage}\n {d}model:{r} {model}\n {d}tokens:{r} {} ({} in / {} out)\n {d}tool_calls:{r} {tool_call_count}\n {d}text:{r} {truncated}\n",
|
||||
format_tokens_human(total),
|
||||
format_tokens_human(*input_tokens),
|
||||
format_tokens_human(*output_tokens),
|
||||
)
|
||||
}
|
||||
PipelineEvent::ToolCallStarted {
|
||||
stage,
|
||||
tool_name,
|
||||
tool_call_id,
|
||||
arguments,
|
||||
} => {
|
||||
let args_str = serde_json::to_string(arguments).unwrap_or_else(|_| arguments.to_string());
|
||||
let truncated = if args_str.len() > 200 { &args_str[..200] } else { &args_str };
|
||||
format!("{d}── TOOL_CALL_STARTED ────────────────────────{r}\n {d}stage:{r} {stage}\n {d}tool_name:{r} {tool_name}\n {d}tool_call_id:{r} {tool_call_id}\n {d}arguments:{r} {truncated}\n")
|
||||
}
|
||||
PipelineEvent::ToolCallCompleted {
|
||||
stage,
|
||||
tool_name,
|
||||
tool_call_id,
|
||||
output,
|
||||
is_error,
|
||||
} => {
|
||||
let output_str = serde_json::to_string(output).unwrap_or_else(|_| output.to_string());
|
||||
let truncated = if output_str.len() > 200 { &output_str[..200] } else { &output_str };
|
||||
format!("{d}── TOOL_CALL_COMPLETED ──────────────────────{r}\n {d}stage:{r} {stage}\n {d}tool_name:{r} {tool_name}\n {d}tool_call_id:{r} {tool_call_id}\n {d}is_error:{r} {is_error}\n {d}output:{r} {truncated}\n")
|
||||
}
|
||||
PipelineEvent::SessionError { stage, error } => {
|
||||
format!("{d}── SESSION_ERROR ────────────────────────────{r}\n {d}stage:{r} {stage}\n {d}error:{r} {error}\n")
|
||||
}
|
||||
PipelineEvent::ContextWindowWarning {
|
||||
stage,
|
||||
estimated_tokens,
|
||||
context_window_size,
|
||||
usage_percent,
|
||||
} => {
|
||||
format!("{d}── CONTEXT_WINDOW_WARNING ───────────────────{r}\n {d}stage:{r} {stage}\n {d}estimated_tokens:{r} {estimated_tokens}\n {d}context_window_size:{r} {context_window_size}\n {d}usage_percent:{r} {usage_percent}%\n")
|
||||
}
|
||||
PipelineEvent::LoopDetected { stage } => {
|
||||
format!("{d}── LOOP_DETECTED ────────────────────────────{r}\n {d}stage:{r} {stage}\n")
|
||||
}
|
||||
PipelineEvent::TurnLimitReached { stage } => {
|
||||
format!("{d}── TURN_LIMIT_REACHED ───────────────────────{r}\n {d}stage:{r} {stage}\n")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -77,6 +77,47 @@ pub enum PipelineEvent {
|
|||
CheckpointSaved {
|
||||
node_id: String,
|
||||
},
|
||||
Prompt {
|
||||
stage: String,
|
||||
text: String,
|
||||
},
|
||||
AssistantMessage {
|
||||
stage: String,
|
||||
text: String,
|
||||
model: String,
|
||||
input_tokens: i64,
|
||||
output_tokens: i64,
|
||||
tool_call_count: usize,
|
||||
},
|
||||
ToolCallStarted {
|
||||
stage: String,
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
arguments: serde_json::Value,
|
||||
},
|
||||
ToolCallCompleted {
|
||||
stage: String,
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
output: serde_json::Value,
|
||||
is_error: bool,
|
||||
},
|
||||
SessionError {
|
||||
stage: String,
|
||||
error: String,
|
||||
},
|
||||
ContextWindowWarning {
|
||||
stage: String,
|
||||
estimated_tokens: usize,
|
||||
context_window_size: usize,
|
||||
usage_percent: usize,
|
||||
},
|
||||
LoopDetected {
|
||||
stage: String,
|
||||
},
|
||||
TurnLimitReached {
|
||||
stage: String,
|
||||
},
|
||||
}
|
||||
|
||||
/// Listener callback type for pipeline events.
|
||||
|
|
@ -168,4 +209,37 @@ mod tests {
|
|||
let emitter = EventEmitter::default();
|
||||
assert_eq!(emitter.listeners.len(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn llm_conversation_event_serialization() {
|
||||
let event = PipelineEvent::ToolCallStarted {
|
||||
stage: "plan".to_string(),
|
||||
tool_name: "read_file".to_string(),
|
||||
tool_call_id: "call_1".to_string(),
|
||||
arguments: serde_json::json!({"path": "/tmp/test.txt"}),
|
||||
};
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
assert!(json.contains("ToolCallStarted"));
|
||||
assert!(json.contains("read_file"));
|
||||
assert!(json.contains("plan"));
|
||||
|
||||
// Verify round-trip
|
||||
let deserialized: PipelineEvent = serde_json::from_str(&json).unwrap();
|
||||
assert!(matches!(deserialized, PipelineEvent::ToolCallStarted { stage, .. } if stage == "plan"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn assistant_message_event_serialization() {
|
||||
let event = PipelineEvent::AssistantMessage {
|
||||
stage: "code".to_string(),
|
||||
text: "Here is the implementation".to_string(),
|
||||
model: "claude-opus-4-6".to_string(),
|
||||
input_tokens: 1000,
|
||||
output_tokens: 500,
|
||||
tool_call_count: 3,
|
||||
};
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
assert!(json.contains("AssistantMessage"));
|
||||
assert!(json.contains("claude-opus-4-6"));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
use std::path::Path;
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use crate::context::Context;
|
||||
use crate::error::AttractorError;
|
||||
use crate::event::EventEmitter;
|
||||
use crate::graph::{Graph, Node};
|
||||
use crate::outcome::{Outcome, StageUsage};
|
||||
|
||||
|
|
@ -27,6 +29,7 @@ pub trait CodergenBackend: Send + Sync {
|
|||
prompt: &str,
|
||||
context: &Context,
|
||||
thread_id: Option<&str>,
|
||||
emitter: &Arc<EventEmitter>,
|
||||
) -> Result<CodergenResult, AttractorError>;
|
||||
}
|
||||
|
||||
|
|
@ -175,7 +178,7 @@ impl Handler for CodergenHandler {
|
|||
context: &Context,
|
||||
graph: &Graph,
|
||||
logs_root: &Path,
|
||||
_services: &EngineServices,
|
||||
services: &EngineServices,
|
||||
) -> Result<Outcome, AttractorError> {
|
||||
// 1. Build prompt
|
||||
let raw_prompt = node
|
||||
|
|
@ -203,7 +206,7 @@ impl Handler for CodergenHandler {
|
|||
.get("internal.thread_id")
|
||||
.and_then(|v| v.as_str().map(String::from));
|
||||
let (response_text, stage_usage) = if let Some(backend) = &self.backend {
|
||||
match backend.run(node, &prompt, context, thread_id.as_deref()).await {
|
||||
match backend.run(node, &prompt, context, thread_id.as_deref(), &services.emitter).await {
|
||||
Ok(CodergenResult::Full(outcome)) => {
|
||||
let status_json = serde_json::to_string_pretty(&outcome)
|
||||
.unwrap_or_else(|_| "{}".to_string());
|
||||
|
|
@ -521,6 +524,7 @@ mod tests {
|
|||
_prompt: &str,
|
||||
_context: &Context,
|
||||
thread_id: Option<&str>,
|
||||
_emitter: &Arc<EventEmitter>,
|
||||
) -> Result<CodergenResult, AttractorError> {
|
||||
*self.captured_thread_id.lock().unwrap() =
|
||||
Some(thread_id.map(String::from));
|
||||
|
|
@ -566,6 +570,7 @@ mod tests {
|
|||
_prompt: &str,
|
||||
_context: &Context,
|
||||
thread_id: Option<&str>,
|
||||
_emitter: &Arc<EventEmitter>,
|
||||
) -> Result<CodergenResult, AttractorError> {
|
||||
*self.captured_thread_id.lock().unwrap() =
|
||||
Some(thread_id.map(String::from));
|
||||
|
|
@ -606,6 +611,7 @@ mod tests {
|
|||
_prompt: &str,
|
||||
_context: &Context,
|
||||
_thread_id: Option<&str>,
|
||||
_emitter: &Arc<EventEmitter>,
|
||||
) -> Result<CodergenResult, AttractorError> {
|
||||
Err(AttractorError::Handler("Request timed out".to_string()))
|
||||
}
|
||||
|
|
@ -717,6 +723,7 @@ Some text in between.
|
|||
_prompt: &str,
|
||||
_context: &Context,
|
||||
_thread_id: Option<&str>,
|
||||
_emitter: &Arc<EventEmitter>,
|
||||
) -> Result<CodergenResult, AttractorError> {
|
||||
Err(AttractorError::Validation("bad config".to_string()))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
use std::path::Path;
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use crate::context::Context;
|
||||
use crate::error::AttractorError;
|
||||
use crate::event::EventEmitter;
|
||||
use crate::graph::{Graph, Node};
|
||||
use crate::outcome::Outcome;
|
||||
|
||||
|
|
@ -30,7 +32,7 @@ impl Handler for FanInHandler {
|
|||
context: &Context,
|
||||
_graph: &Graph,
|
||||
logs_root: &Path,
|
||||
_services: &EngineServices,
|
||||
services: &EngineServices,
|
||||
) -> Result<Outcome, AttractorError> {
|
||||
let results = context.get("parallel.results");
|
||||
let Some(results) = results else {
|
||||
|
|
@ -40,7 +42,7 @@ impl Handler for FanInHandler {
|
|||
let prompt = node.prompt().filter(|p| !p.is_empty());
|
||||
|
||||
let best = if let (Some(prompt_text), Some(backend)) = (prompt, &self.backend) {
|
||||
llm_evaluate(backend.as_ref(), prompt_text, &results, context, logs_root, &node.id).await?
|
||||
llm_evaluate(backend.as_ref(), prompt_text, &results, context, logs_root, &node.id, &services.emitter).await?
|
||||
} else {
|
||||
heuristic_select(&results)
|
||||
};
|
||||
|
|
@ -153,6 +155,7 @@ async fn llm_evaluate(
|
|||
context: &Context,
|
||||
logs_root: &Path,
|
||||
node_id: &str,
|
||||
emitter: &Arc<EventEmitter>,
|
||||
) -> Result<Candidate, AttractorError> {
|
||||
let results_text = serde_json::to_string_pretty(results)
|
||||
.unwrap_or_else(|_| results.to_string());
|
||||
|
|
@ -171,7 +174,7 @@ async fn llm_evaluate(
|
|||
let eval_node = Node::new("fan_in_eval");
|
||||
|
||||
// Fan-in evaluation runs outside a thread context, so pass None
|
||||
match backend.run(&eval_node, &full_prompt, context, None).await {
|
||||
match backend.run(&eval_node, &full_prompt, context, None, emitter).await {
|
||||
Ok(CodergenResult::Full(outcome)) => {
|
||||
// If the backend returned a full Outcome, extract best_id from context_updates
|
||||
let best_id = outcome
|
||||
|
|
@ -367,6 +370,7 @@ mod tests {
|
|||
_prompt: &str,
|
||||
_context: &Context,
|
||||
_thread_id: Option<&str>,
|
||||
_emitter: &Arc<EventEmitter>,
|
||||
) -> Result<CodergenResult, AttractorError> {
|
||||
// Return text that contains the ID "branch_b"
|
||||
Ok(CodergenResult::Text { text: "The best candidate is branch_b".to_string(), usage: None })
|
||||
|
|
|
|||
|
|
@ -1056,6 +1056,7 @@ impl CodergenBackend for MockCodergenBackend {
|
|||
prompt: &str,
|
||||
_context: &Context,
|
||||
_thread_id: Option<&str>,
|
||||
_emitter: &Arc<EventEmitter>,
|
||||
) -> Result<CodergenResult, AttractorError> {
|
||||
Ok(CodergenResult::Text {
|
||||
text: format!(
|
||||
|
|
@ -5079,6 +5080,7 @@ mod real_llm {
|
|||
prompt: &str,
|
||||
_context: &Context,
|
||||
_thread_id: Option<&str>,
|
||||
_emitter: &Arc<EventEmitter>,
|
||||
) -> Result<CodergenResult, AttractorError> {
|
||||
let request = Request {
|
||||
model: self.model.clone(),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue