1use crate::capabilities::CapabilityRegistry;
15use crate::command::CommandResult;
16use crate::driver_registry::{
17 BoxedChatDriver, DriverRegistry, LlmCallConfig, LlmMessage, LlmMessageRole, LlmResponseStream,
18 ToolSearchConfig,
19};
20use crate::error::{AgentLoopError, Result};
21use crate::message::{Controls, Message, MessageRole, patch_dangling_tool_calls};
22use crate::message_retriever::MessageRetriever;
23use crate::runtime_context::{AssembledTurnContext, inspect_turn_context};
24use crate::session::Session;
25use crate::traits::{
26 AgentStore, HarnessStore, ImageResolver, ProviderStore, ResolvedImage, ResolvedModel,
27 SessionFileSystem, SessionStore,
28};
29use crate::typed_id::SessionId;
30use crate::user_facing_error::{UserFacingErrorContext, classify_runtime_error_message};
31use async_trait::async_trait;
32use std::collections::{HashMap, HashSet};
33use std::sync::Arc;
34use uuid::Uuid;
35
36#[derive(Debug, Clone)]
45pub struct CommandTurnContext {
46 pub session: Session,
48 pub messages: Vec<Message>,
50 pub system_prompt: String,
52 pub model: String,
54 pub provider_type: String,
56 pub resolved_locale: Option<String>,
58}
59
60#[derive(Debug, Clone, Default)]
65pub struct SessionCompletionRequest {
66 pub system_prompts: Vec<String>,
69 pub messages: Vec<Message>,
71 pub controls: Option<Controls>,
74 pub metadata: HashMap<String, String>,
77}
78
79#[derive(Debug, Clone)]
81pub struct SessionCompletion {
82 pub text: String,
84}
85
86pub struct SessionCompletionStream {
90 pub events: LlmResponseStream,
92 pub context: UserFacingErrorContext,
96}
97
98impl std::fmt::Debug for SessionCompletionStream {
99 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
100 f.debug_struct("SessionCompletionStream")
101 .field("context", &self.context)
102 .finish()
103 }
104}
105
106#[derive(Debug)]
108pub enum SessionCompletionError {
109 InvalidRequest(AgentLoopError),
112 StreamingUnsupported,
116 Completion {
119 error: String,
121 context: UserFacingErrorContext,
123 },
124}
125
126impl SessionCompletionError {
127 pub fn into_command_result(self) -> Result<CommandResult> {
131 match self {
132 Self::InvalidRequest(error) => Err(error),
133 Self::StreamingUnsupported => Err(AgentLoopError::config(
134 "command host does not support streaming completions",
135 )),
136 Self::Completion { error, context } => {
137 let classified = classify_runtime_error_message(&error, &context);
138 Ok(CommandResult {
139 success: false,
140 message: classified.fallback_message(),
141 error_code: Some(classified.code.clone()),
142 error_fields: classified.error_fields(),
143 })
144 }
145 }
146 }
147}
148
149#[async_trait]
154pub trait CommandHost: Send + Sync {
155 async fn turn_context(&self) -> Result<CommandTurnContext>;
157
158 async fn completion(
161 &self,
162 request: SessionCompletionRequest,
163 ) -> std::result::Result<SessionCompletion, SessionCompletionError>;
164
165 async fn completion_stream(
173 &self,
174 _request: SessionCompletionRequest,
175 ) -> std::result::Result<SessionCompletionStream, SessionCompletionError> {
176 Err(SessionCompletionError::StreamingUnsupported)
177 }
178}
179
180pub struct DisabledCommandHost;
183
184#[async_trait]
185impl CommandHost for DisabledCommandHost {
186 async fn turn_context(&self) -> Result<CommandTurnContext> {
187 Err(AgentLoopError::config(
188 "command host does not provide turn-context access",
189 ))
190 }
191
192 async fn completion(
193 &self,
194 _request: SessionCompletionRequest,
195 ) -> std::result::Result<SessionCompletion, SessionCompletionError> {
196 Err(SessionCompletionError::InvalidRequest(
197 AgentLoopError::config("command host does not provide session completions"),
198 ))
199 }
200}
201
202pub struct StoreCommandHost {
210 session_id: SessionId,
211 harness_store: Arc<dyn HarnessStore>,
212 agent_store: Arc<dyn AgentStore>,
213 session_store: Arc<dyn SessionStore>,
214 message_retriever: Arc<dyn MessageRetriever>,
215 provider_store: Arc<dyn ProviderStore>,
216 capability_registry: CapabilityRegistry,
217 driver_registry: DriverRegistry,
218 image_resolver: Option<Arc<dyn ImageResolver>>,
219 file_store: Option<Arc<dyn SessionFileSystem>>,
220 assembled: tokio::sync::OnceCell<AssembledTurnContext>,
221}
222
223impl StoreCommandHost {
224 #[allow(clippy::too_many_arguments)]
225 pub fn new(
226 session_id: SessionId,
227 harness_store: Arc<dyn HarnessStore>,
228 agent_store: Arc<dyn AgentStore>,
229 session_store: Arc<dyn SessionStore>,
230 message_retriever: Arc<dyn MessageRetriever>,
231 provider_store: Arc<dyn ProviderStore>,
232 capability_registry: CapabilityRegistry,
233 driver_registry: DriverRegistry,
234 ) -> Self {
235 Self {
236 session_id,
237 harness_store,
238 agent_store,
239 session_store,
240 message_retriever,
241 provider_store,
242 capability_registry,
243 driver_registry,
244 image_resolver: None,
245 file_store: None,
246 assembled: tokio::sync::OnceCell::new(),
247 }
248 }
249
250 pub fn with_image_resolver(mut self, image_resolver: Arc<dyn ImageResolver>) -> Self {
253 self.image_resolver = Some(image_resolver);
254 self
255 }
256
257 pub fn with_file_store(mut self, file_store: Arc<dyn SessionFileSystem>) -> Self {
260 self.file_store = Some(file_store);
261 self
262 }
263
264 pub fn with_assembled_context(mut self, assembled: AssembledTurnContext) -> Self {
267 self.assembled = tokio::sync::OnceCell::new_with(Some(assembled));
268 self
269 }
270
271 async fn assembled(&self) -> Result<&AssembledTurnContext> {
272 self.assembled
273 .get_or_try_init(|| async {
274 let session = self
275 .session_store
276 .get_session(self.session_id)
277 .await?
278 .ok_or_else(|| AgentLoopError::session_not_found(self.session_id))?;
279 inspect_turn_context(
280 self.harness_store.as_ref(),
281 self.agent_store.as_ref(),
282 self.session_store.as_ref(),
283 self.message_retriever.as_ref(),
284 self.provider_store.as_ref(),
285 &self.capability_registry,
286 self.session_id,
287 session.harness_id,
288 session.agent_id,
289 &[],
290 self.file_store.clone(),
291 )
292 .await
293 })
294 .await
295 }
296
297 async fn resolve_images(&self, messages: &[Message]) -> HashMap<Uuid, ResolvedImage> {
298 let Some(resolver) = &self.image_resolver else {
299 return HashMap::new();
300 };
301 let image_ids: HashSet<Uuid> = messages
302 .iter()
303 .flat_map(crate::llm_conversions::extract_image_file_ids)
304 .collect();
305 let mut resolved = HashMap::new();
306 for image_id in image_ids {
307 if let Ok(Some(image)) = resolver.resolve_image(image_id).await {
308 resolved.insert(image_id, image);
309 }
310 }
311 resolved
312 }
313
314 async fn resolve_completion_model(
318 &self,
319 controls: Option<&Controls>,
320 assembled: &AssembledTurnContext,
321 ) -> std::result::Result<ResolvedModel, SessionCompletionError> {
322 let requested = controls.and_then(|controls| controls.model_id);
323 match requested {
324 Some(model_id) if Some(model_id) != assembled.resolved_model_id => self
325 .provider_store
326 .get_resolved_model(model_id)
327 .await
328 .map_err(SessionCompletionError::InvalidRequest)?
329 .ok_or_else(|| {
330 SessionCompletionError::InvalidRequest(AgentLoopError::config(format!(
331 "Model not found: {model_id}"
332 )))
333 }),
334 _ => Ok(assembled.model_with_provider.clone()),
335 }
336 }
337
338 async fn prepare_completion(
342 &self,
343 request: SessionCompletionRequest,
344 ) -> std::result::Result<PreparedCompletion, SessionCompletionError> {
345 let assembled = self
346 .assembled()
347 .await
348 .map_err(SessionCompletionError::InvalidRequest)?;
349 let model = self
350 .resolve_completion_model(request.controls.as_ref(), assembled)
351 .await?;
352
353 let context = UserFacingErrorContext::default()
354 .with_provider(model.provider_type.to_string())
355 .with_model_id(model.model.clone());
356
357 let messages = patch_dangling_tool_calls(&request.messages);
358 let resolved_images = self.resolve_images(&messages).await;
359
360 let mut llm_messages: Vec<LlmMessage> = request
361 .system_prompts
362 .iter()
363 .filter(|prompt| !prompt.is_empty())
364 .map(|prompt| LlmMessage::text(LlmMessageRole::System, prompt.clone()))
365 .collect();
366 for msg in &messages {
367 let mut llm_msg =
368 crate::llm_conversions::llm_message_from_message_with_images(msg, &resolved_images);
369 if msg.role == MessageRole::User
370 && let Some(actor) = &msg.external_actor
371 {
372 llm_msg.prepend_text_prefix(&format!("[{}] ", actor.display_label()));
373 }
374 llm_messages.push(llm_msg);
375 }
376
377 let mut llm_config_builder =
378 crate::llm_conversions::llm_call_config_builder_from_agent(&assembled.runtime_agent)
379 .model(&model.model)
380 .tools(vec![])
381 .tool_search(ToolSearchConfig {
382 enabled: false,
383 threshold: usize::MAX,
384 })
385 .previous_response_id(None)
386 .with_metadata("session_id", self.session_id.to_string());
387 if let Some(effort) = request
388 .controls
389 .as_ref()
390 .and_then(|controls| controls.reasoning.as_ref())
391 .and_then(|reasoning| reasoning.effort.clone())
392 .filter(|value| !value.is_empty())
393 {
394 llm_config_builder = llm_config_builder.reasoning_effort(effort);
395 }
396 for (key, value) in &request.metadata {
397 llm_config_builder = llm_config_builder.with_metadata(key, value);
398 }
399 let llm_config = llm_config_builder.build();
400
401 let driver = self
402 .driver_registry
403 .create_chat_driver(
404 &crate::llm_conversions::provider_config_from_resolved_model(&model),
405 )
406 .map_err(|error| SessionCompletionError::Completion {
407 error: error.to_string(),
408 context: context.clone(),
409 })?;
410
411 Ok(PreparedCompletion {
412 llm_messages,
413 llm_config,
414 driver,
415 context,
416 })
417 }
418}
419
420struct PreparedCompletion {
423 llm_messages: Vec<LlmMessage>,
424 llm_config: LlmCallConfig,
425 driver: BoxedChatDriver,
426 context: UserFacingErrorContext,
427}
428
429#[async_trait]
430impl CommandHost for StoreCommandHost {
431 async fn turn_context(&self) -> Result<CommandTurnContext> {
432 let assembled = self.assembled().await?;
433 Ok(CommandTurnContext {
434 session: assembled.session.clone(),
435 messages: assembled.messages.clone(),
436 system_prompt: assembled.runtime_agent.system_prompt.clone(),
437 model: assembled.model_with_provider.model.clone(),
438 provider_type: assembled.model_with_provider.provider_type.to_string(),
439 resolved_locale: assembled.resolved_locale.clone(),
440 })
441 }
442
443 async fn completion(
444 &self,
445 request: SessionCompletionRequest,
446 ) -> std::result::Result<SessionCompletion, SessionCompletionError> {
447 let prepared = self.prepare_completion(request).await?;
448 let completion_error = |error: String| SessionCompletionError::Completion {
449 error,
450 context: prepared.context.clone(),
451 };
452
453 let response = prepared
454 .driver
455 .chat_completion(prepared.llm_messages, &prepared.llm_config)
456 .await
457 .map_err(|error| completion_error(error.to_string()))?;
458
459 let text = response.text.trim().to_string();
460 if text.is_empty() {
461 return Err(completion_error(
462 "session completion returned an empty response".to_string(),
463 ));
464 }
465 Ok(SessionCompletion { text })
466 }
467
468 async fn completion_stream(
469 &self,
470 request: SessionCompletionRequest,
471 ) -> std::result::Result<SessionCompletionStream, SessionCompletionError> {
472 let prepared = self.prepare_completion(request).await?;
473 let events = prepared
474 .driver
475 .chat_completion_stream(prepared.llm_messages, &prepared.llm_config)
476 .await
477 .map_err(|error| SessionCompletionError::Completion {
478 error: error.to_string(),
479 context: prepared.context.clone(),
480 })?;
481 Ok(SessionCompletionStream {
482 events,
483 context: prepared.context,
484 })
485 }
486}
487
488#[cfg(all(test, feature = "llmsim"))]
491mod tests {
492 use super::*;
493 use crate::agent::{Agent, AgentStatus};
494 use crate::capabilities::TestMathCapability;
495 use crate::driver_registry::LlmStreamEvent;
496 use crate::harness::{Harness, HarnessStatus};
497 use crate::in_memory::{
498 InMemoryAgentStore, InMemoryHarnessStore, InMemoryMessageRetriever, InMemoryProviderStore,
499 InMemorySessionStore,
500 };
501 use crate::llmsim_driver::{LlmSimConfig, LlmSimDriver};
502 use crate::message_retriever::InputMessage;
503 use crate::provider::DriverId;
504 use crate::session::SessionStatus;
505 use crate::typed_id::{AgentId, HarnessId};
506 use chrono::Utc;
507 use futures::StreamExt;
508
509 #[tokio::test]
510 async fn disabled_host_errors_clearly() {
511 let host = DisabledCommandHost;
512 let error = host.turn_context().await.unwrap_err();
513 assert!(error.to_string().contains("turn-context"));
514
515 let error = host
516 .completion(SessionCompletionRequest::default())
517 .await
518 .unwrap_err();
519 assert!(matches!(error, SessionCompletionError::InvalidRequest(_)));
520
521 let error = host
524 .completion_stream(SessionCompletionRequest::default())
525 .await
526 .unwrap_err();
527 assert!(matches!(
528 error,
529 SessionCompletionError::StreamingUnsupported
530 ));
531 let error = error.into_command_result().unwrap_err();
532 assert!(error.to_string().contains("streaming"));
533 }
534
535 fn test_harness(harness_id: HarnessId) -> Harness {
536 Harness {
537 id: harness_id,
538 name: "h".into(),
539 display_name: None,
540 description: None,
541 system_prompt: Some("You are a test harness.".into()),
542 parent_harness_id: None,
543 default_model_id: None,
544 tags: vec![],
545 capabilities: vec![crate::AgentCapabilityConfig::new("test_math")],
546 initial_files: vec![],
547 network_access: None,
548 parallel_tool_calls: None,
549 mcp_servers: Default::default(),
550 embedder_metadata: Default::default(),
551 is_built_in: false,
552 status: HarnessStatus::Active,
553 created_at: Utc::now(),
554 updated_at: Utc::now(),
555 archived_at: None,
556 deleted_at: None,
557 }
558 }
559
560 fn test_agent(agent_id: AgentId) -> Agent {
561 Agent {
562 public_id: agent_id,
563 internal_id: uuid::Uuid::nil(),
564 name: "a".into(),
565 display_name: None,
566 description: None,
567 system_prompt: "Use tools.".into(),
568 default_model_id: None,
569
570 harness_id: crate::typed_id::HarnessId::from_uuid(uuid::Uuid::nil()),
571 default_version_id: None,
572 forked_from_agent_id: None,
573 forked_from_version_id: None,
574 root_agent_id: None,
575 tags: vec![],
576 capabilities: vec![],
577 initial_files: vec![],
578 network_access: None,
579 max_iterations: Some(8),
580 parallel_tool_calls: None,
581 tools: vec![],
582 mcp_servers: Default::default(),
583 status: AgentStatus::Active,
584 created_at: Utc::now(),
585 updated_at: Utc::now(),
586 archived_at: None,
587 deleted_at: None,
588 usage: None,
589 }
590 }
591
592 fn test_session(session_id: SessionId, harness_id: HarnessId, agent_id: AgentId) -> Session {
593 Session {
594 id: session_id,
595 workspace_id: crate::WorkspaceId::from_uuid((session_id).uuid()),
596 organization_id: crate::DEFAULT_ORG_PUBLIC_ID.to_string(),
597 harness_id,
598 agent_id: Some(agent_id),
599 agent_version_id: None,
600 agent_identity_id: None,
601 owner_principal_id: crate::PrincipalId::from_seed(1),
602 resolved_owner_user_id: None,
603 owner: None,
604 effective_owner: None,
605 title: None,
606 goal: None,
607 locale: None,
608 preview: None,
609 output_preview: None,
610 tags: vec![],
611 model_id: None,
612 capabilities: vec![],
613 tools: vec![],
614 mcp_servers: Default::default(),
615 system_prompt: None,
616 initial_files: vec![],
617 hints: None,
618 network_access: None,
619 max_iterations: None,
620 parallel_tool_calls: None,
621 status: SessionStatus::Started,
622 created_at: Utc::now(),
623 updated_at: Utc::now(),
624 started_at: None,
625 finished_at: None,
626 usage: None,
627 is_pinned: None,
628 active_schedule_count: None,
629 features: vec![],
630 parent_session_id: None,
631 forked_from_session_id: None,
632 forked_from_sequence: None,
633 blueprint_id: None,
634 blueprint_config: None,
635 }
636 }
637
638 async fn llmsim_host(response: &str) -> StoreCommandHost {
641 let harness_id: HarnessId = "harness_000000000000000000000000000000a1".parse().unwrap();
642 let agent_id: AgentId = "agent_000000000000000000000000000000a1".parse().unwrap();
643 let session_id: SessionId = "session_000000000000000000000000000000a1".parse().unwrap();
644
645 let harness_store = InMemoryHarnessStore::new();
646 harness_store.add_harness(test_harness(harness_id)).await;
647 let agent_store = InMemoryAgentStore::new();
648 agent_store.add_agent(test_agent(agent_id)).await;
649 let session_store = InMemorySessionStore::new();
650 session_store
651 .add_session(test_session(session_id, harness_id, agent_id))
652 .await;
653 let message_store = InMemoryMessageRetriever::new();
654 message_store
655 .add(session_id, InputMessage::user("earlier message"))
656 .await
657 .unwrap();
658
659 let provider_store = InMemoryProviderStore::new();
660 provider_store
661 .set_default_model(ResolvedModel {
662 model: "llmsim-model".into(),
663 provider_type: DriverId::LlmSim,
664 api_key: Some("fake-key".into()),
665 base_url: None,
666 provider_metadata: None,
667 })
668 .await;
669
670 let mut capability_registry = CapabilityRegistry::new();
671 capability_registry.register(TestMathCapability);
672
673 let mut driver_registry = DriverRegistry::new();
674 let driver = LlmSimDriver::new(LlmSimConfig::fixed(response));
675 driver_registry.register(DriverId::LlmSim, move |_config| Box::new(driver.clone()));
676
677 StoreCommandHost::new(
678 session_id,
679 Arc::new(harness_store),
680 Arc::new(agent_store),
681 Arc::new(session_store),
682 Arc::new(message_store),
683 Arc::new(provider_store),
684 capability_registry,
685 driver_registry,
686 )
687 }
688
689 #[tokio::test]
690 async fn store_host_completion_runs_against_session_model() {
691 let host = llmsim_host("the side answer").await;
692
693 let turn = host.turn_context().await.unwrap();
694 assert_eq!(turn.model, "llmsim-model");
695 assert_eq!(turn.provider_type, "llmsim");
696 assert_eq!(turn.messages.len(), 1);
697 assert!(!turn.system_prompt.is_empty());
698
699 let completion = host
700 .completion(SessionCompletionRequest {
701 system_prompts: vec![turn.system_prompt, "Answer once.".into()],
702 messages: turn.messages,
703 controls: None,
704 metadata: HashMap::new(),
705 })
706 .await
707 .unwrap();
708 assert_eq!(completion.text, "the side answer");
709 }
710
711 #[tokio::test]
712 async fn store_host_completion_stream_emits_progressive_deltas() {
713 let host = llmsim_host("streamed side answer with several tokens").await;
714
715 let turn = host.turn_context().await.unwrap();
716 let stream = host
717 .completion_stream(SessionCompletionRequest {
718 system_prompts: vec![turn.system_prompt],
719 messages: turn.messages,
720 controls: None,
721 metadata: HashMap::new(),
722 })
723 .await
724 .unwrap();
725
726 assert_eq!(stream.context.provider.as_deref(), Some("llmsim"));
728 assert_eq!(stream.context.model_id.as_deref(), Some("llmsim-model"));
729
730 let mut deltas = Vec::new();
731 let mut done = false;
732 let mut events = stream.events;
733 while let Some(event) = events.next().await {
734 match event.unwrap() {
735 LlmStreamEvent::TextDelta(delta) => deltas.push(delta),
736 LlmStreamEvent::Done(_) => done = true,
737 _ => {}
738 }
739 }
740
741 assert!(done, "stream must terminate with Done");
742 assert!(
743 deltas.len() > 1,
744 "expected progressive deltas, got {deltas:?}"
745 );
746 assert_eq!(deltas.concat(), "streamed side answer with several tokens");
747 }
748
749 #[test]
750 fn completion_error_classifies_provider_failures() {
751 let error = SessionCompletionError::Completion {
752 error: "OpenAI API error (401): unauthorized".to_string(),
753 context: UserFacingErrorContext::default()
754 .with_provider("openai")
755 .with_model_id("gpt-5"),
756 };
757
758 let result = error.into_command_result().expect("classified result");
759 assert!(!result.success);
760 assert_eq!(result.error_code.as_deref(), Some("provider_misconfigured"));
761 let fields = result.error_fields.expect("error_fields populated");
762 assert_eq!(
763 fields.get("provider").and_then(|v| v.as_str()),
764 Some("openai")
765 );
766 assert_eq!(
767 fields.get("model_id").and_then(|v| v.as_str()),
768 Some("gpt-5")
769 );
770 }
771
772 #[test]
773 fn completion_error_bubbles_invalid_requests() {
774 let error =
775 SessionCompletionError::InvalidRequest(AgentLoopError::config("Model not found"));
776 assert!(error.into_command_result().is_err());
777 }
778}