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 (_, compatibility_config) = model.canonical_parts();
402 let provider_config = self
403 .provider_store
404 .get_provider_config(&compatibility_config.provider)
405 .await
406 .map_err(|error| SessionCompletionError::Completion {
407 error: error.to_string(),
408 context: context.clone(),
409 })?
410 .unwrap_or(compatibility_config);
411 let driver = self
412 .driver_registry
413 .create_chat_driver(&provider_config)
414 .map_err(|error| SessionCompletionError::Completion {
415 error: error.to_string(),
416 context: context.clone(),
417 })?;
418
419 Ok(PreparedCompletion {
420 llm_messages,
421 llm_config,
422 driver,
423 context,
424 })
425 }
426}
427
428struct PreparedCompletion {
431 llm_messages: Vec<LlmMessage>,
432 llm_config: LlmCallConfig,
433 driver: BoxedChatDriver,
434 context: UserFacingErrorContext,
435}
436
437#[async_trait]
438impl CommandHost for StoreCommandHost {
439 async fn turn_context(&self) -> Result<CommandTurnContext> {
440 let assembled = self.assembled().await?;
441 Ok(CommandTurnContext {
442 session: assembled.session.clone(),
443 messages: assembled.messages.clone(),
444 system_prompt: assembled.runtime_agent.system_prompt.clone(),
445 model: assembled.model_with_provider.model.clone(),
446 provider_type: assembled.model_with_provider.provider_type.to_string(),
447 resolved_locale: assembled.resolved_locale.clone(),
448 })
449 }
450
451 async fn completion(
452 &self,
453 request: SessionCompletionRequest,
454 ) -> std::result::Result<SessionCompletion, SessionCompletionError> {
455 let prepared = self.prepare_completion(request).await?;
456 let completion_error = |error: String| SessionCompletionError::Completion {
457 error,
458 context: prepared.context.clone(),
459 };
460
461 let response = prepared
462 .driver
463 .chat_completion(
464 &crate::ProviderEndpoint::default(),
465 prepared.llm_messages,
466 &prepared.llm_config,
467 )
468 .await
469 .map_err(|error| completion_error(error.to_string()))?;
470
471 let text = response.text.trim().to_string();
472 if text.is_empty() {
473 return Err(completion_error(
474 "session completion returned an empty response".to_string(),
475 ));
476 }
477 Ok(SessionCompletion { text })
478 }
479
480 async fn completion_stream(
481 &self,
482 request: SessionCompletionRequest,
483 ) -> std::result::Result<SessionCompletionStream, SessionCompletionError> {
484 let prepared = self.prepare_completion(request).await?;
485 let events = prepared
486 .driver
487 .chat_completion_stream(
488 &crate::ProviderEndpoint::default(),
489 prepared.llm_messages,
490 &prepared.llm_config,
491 )
492 .await
493 .map_err(|error| SessionCompletionError::Completion {
494 error: error.to_string(),
495 context: prepared.context.clone(),
496 })?;
497 Ok(SessionCompletionStream {
498 events,
499 context: prepared.context,
500 })
501 }
502}
503
504#[cfg(all(test, feature = "llmsim"))]
507mod tests {
508 use super::*;
509 use crate::agent::{Agent, AgentStatus};
510 use crate::capabilities::TestMathCapability;
511 use crate::driver_registry::LlmStreamEvent;
512 use crate::harness::{Harness, HarnessStatus};
513 use crate::in_memory::{
514 InMemoryAgentStore, InMemoryHarnessStore, InMemoryMessageRetriever, InMemoryProviderStore,
515 InMemorySessionStore,
516 };
517 use crate::llmsim_driver::{LlmSimConfig, LlmSimDriver};
518 use crate::message_retriever::InputMessage;
519 use crate::provider::DriverId;
520 use crate::session::SessionStatus;
521 use crate::typed_id::{AgentId, HarnessId};
522 use chrono::Utc;
523 use futures::StreamExt;
524
525 #[tokio::test]
526 async fn disabled_host_errors_clearly() {
527 let host = DisabledCommandHost;
528 let error = host.turn_context().await.unwrap_err();
529 assert!(error.to_string().contains("turn-context"));
530
531 let error = host
532 .completion(SessionCompletionRequest::default())
533 .await
534 .unwrap_err();
535 assert!(matches!(error, SessionCompletionError::InvalidRequest(_)));
536
537 let error = host
540 .completion_stream(SessionCompletionRequest::default())
541 .await
542 .unwrap_err();
543 assert!(matches!(
544 error,
545 SessionCompletionError::StreamingUnsupported
546 ));
547 let error = error.into_command_result().unwrap_err();
548 assert!(error.to_string().contains("streaming"));
549 }
550
551 fn test_harness(harness_id: HarnessId) -> Harness {
552 Harness {
553 id: harness_id,
554 name: "h".into(),
555 display_name: None,
556 description: None,
557 system_prompt: Some("You are a test harness.".into()),
558 parent_harness_id: None,
559 default_model_id: None,
560 tags: vec![],
561 capabilities: vec![crate::AgentCapabilityConfig::new("test_math")],
562 initial_files: vec![],
563 network_access: None,
564 parallel_tool_calls: None,
565 mcp_servers: Default::default(),
566 embedder_metadata: Default::default(),
567 is_built_in: false,
568 status: HarnessStatus::Active,
569 created_at: Utc::now(),
570 updated_at: Utc::now(),
571 archived_at: None,
572 deleted_at: None,
573 }
574 }
575
576 fn test_agent(agent_id: AgentId) -> Agent {
577 Agent {
578 public_id: agent_id,
579 internal_id: uuid::Uuid::nil(),
580 name: "a".into(),
581 display_name: None,
582 description: None,
583 system_prompt: "Use tools.".into(),
584 default_model_id: None,
585
586 harness_id: crate::typed_id::HarnessId::from_uuid(uuid::Uuid::nil()),
587 default_version_id: None,
588 forked_from_agent_id: None,
589 forked_from_version_id: None,
590 root_agent_id: None,
591 tags: vec![],
592 capabilities: vec![],
593 initial_files: vec![],
594 network_access: None,
595 max_iterations: Some(8),
596 parallel_tool_calls: None,
597 tools: vec![],
598 mcp_servers: Default::default(),
599 status: AgentStatus::Active,
600 created_at: Utc::now(),
601 updated_at: Utc::now(),
602 archived_at: None,
603 deleted_at: None,
604 usage: None,
605 }
606 }
607
608 fn test_session(session_id: SessionId, harness_id: HarnessId, agent_id: AgentId) -> Session {
609 Session {
610 source: Default::default(),
611 activity: Default::default(),
612 id: session_id,
613 workspace_id: crate::WorkspaceId::from_uuid((session_id).uuid()),
614 organization_id: crate::DEFAULT_ORG_PUBLIC_ID.to_string(),
615 harness_id,
616 agent_id: Some(agent_id),
617 agent_version_id: None,
618 agent_identity_id: None,
619 owner_principal_id: crate::PrincipalId::from_seed(1),
620 resolved_owner_user_id: None,
621 owner: None,
622 effective_owner: None,
623 title: None,
624 goal: None,
625 locale: None,
626 preview: None,
627 output_preview: None,
628 tags: vec![],
629 model_id: None,
630 capabilities: vec![],
631 tools: vec![],
632 mcp_servers: Default::default(),
633 system_prompt: None,
634 initial_files: vec![],
635 hints: None,
636 network_access: None,
637 max_iterations: None,
638 parallel_tool_calls: None,
639 status: SessionStatus::Started,
640 created_at: Utc::now(),
641 updated_at: Utc::now(),
642 started_at: None,
643 finished_at: None,
644 usage: None,
645 is_pinned: None,
646 active_schedule_count: None,
647 features: vec![],
648 parent_session_id: None,
649 forked_from_session_id: None,
650 forked_from_sequence: None,
651 blueprint_id: None,
652 blueprint_config: None,
653 }
654 }
655
656 async fn llmsim_host(response: &str) -> StoreCommandHost {
659 let harness_id: HarnessId = "harness_000000000000000000000000000000a1".parse().unwrap();
660 let agent_id: AgentId = "agent_000000000000000000000000000000a1".parse().unwrap();
661 let session_id: SessionId = "session_000000000000000000000000000000a1".parse().unwrap();
662
663 let harness_store = InMemoryHarnessStore::new();
664 harness_store.add_harness(test_harness(harness_id)).await;
665 let agent_store = InMemoryAgentStore::new();
666 agent_store.add_agent(test_agent(agent_id)).await;
667 let session_store = InMemorySessionStore::new();
668 session_store
669 .add_session(test_session(session_id, harness_id, agent_id))
670 .await;
671 let message_store = InMemoryMessageRetriever::new();
672 message_store
673 .add(session_id, InputMessage::user("earlier message"))
674 .await
675 .unwrap();
676
677 let provider_store = InMemoryProviderStore::new();
678 provider_store
679 .set_default_model(ResolvedModel {
680 model: "llmsim-model".into(),
681 provider_type: DriverId::LlmSim,
682 api_key: Some("fake-key".into()),
683 base_url: None,
684 provider_metadata: None,
685 })
686 .await;
687
688 let mut capability_registry = CapabilityRegistry::new();
689 capability_registry.register(TestMathCapability);
690
691 let mut driver_registry = DriverRegistry::new();
692 let driver = LlmSimDriver::new(LlmSimConfig::fixed(response));
693 driver_registry.register(DriverId::LlmSim, move |_config| Box::new(driver.clone()));
694
695 StoreCommandHost::new(
696 session_id,
697 Arc::new(harness_store),
698 Arc::new(agent_store),
699 Arc::new(session_store),
700 Arc::new(message_store),
701 Arc::new(provider_store),
702 capability_registry,
703 driver_registry,
704 )
705 }
706
707 #[tokio::test]
708 async fn store_host_completion_runs_against_session_model() {
709 let host = llmsim_host("the side answer").await;
710
711 let turn = host.turn_context().await.unwrap();
712 assert_eq!(turn.model, "llmsim-model");
713 assert_eq!(turn.provider_type, "llmsim");
714 assert_eq!(turn.messages.len(), 1);
715 assert!(!turn.system_prompt.is_empty());
716
717 let completion = host
718 .completion(SessionCompletionRequest {
719 system_prompts: vec![turn.system_prompt, "Answer once.".into()],
720 messages: turn.messages,
721 controls: None,
722 metadata: HashMap::new(),
723 })
724 .await
725 .unwrap();
726 assert_eq!(completion.text, "the side answer");
727 }
728
729 #[tokio::test]
730 async fn store_host_completion_stream_emits_progressive_deltas() {
731 let host = llmsim_host("streamed side answer with several tokens").await;
732
733 let turn = host.turn_context().await.unwrap();
734 let stream = host
735 .completion_stream(SessionCompletionRequest {
736 system_prompts: vec![turn.system_prompt],
737 messages: turn.messages,
738 controls: None,
739 metadata: HashMap::new(),
740 })
741 .await
742 .unwrap();
743
744 assert_eq!(stream.context.provider.as_deref(), Some("llmsim"));
746 assert_eq!(stream.context.model_id.as_deref(), Some("llmsim-model"));
747
748 let mut deltas = Vec::new();
749 let mut done = false;
750 let mut events = stream.events;
751 while let Some(event) = events.next().await {
752 match event.unwrap() {
753 LlmStreamEvent::TextDelta(delta) => deltas.push(delta),
754 LlmStreamEvent::Done(_) => done = true,
755 _ => {}
756 }
757 }
758
759 assert!(done, "stream must terminate with Done");
760 assert!(
761 deltas.len() > 1,
762 "expected progressive deltas, got {deltas:?}"
763 );
764 assert_eq!(deltas.concat(), "streamed side answer with several tokens");
765 }
766
767 #[test]
768 fn completion_error_classifies_provider_failures() {
769 let error = SessionCompletionError::Completion {
770 error: "OpenAI API error (401): unauthorized".to_string(),
771 context: UserFacingErrorContext::default()
772 .with_provider("openai")
773 .with_model_id("gpt-5"),
774 };
775
776 let result = error.into_command_result().expect("classified result");
777 assert!(!result.success);
778 assert_eq!(result.error_code.as_deref(), Some("provider_misconfigured"));
779 let fields = result.error_fields.expect("error_fields populated");
780 assert_eq!(
781 fields.get("provider").and_then(|v| v.as_str()),
782 Some("openai")
783 );
784 assert_eq!(
785 fields.get("model_id").and_then(|v| v.as_str()),
786 Some("gpt-5")
787 );
788 }
789
790 #[test]
791 fn completion_error_bubbles_invalid_requests() {
792 let error =
793 SessionCompletionError::InvalidRequest(AgentLoopError::config("Model not found"));
794 assert!(error.into_command_result().is_err());
795 }
796}