Skip to main content

everruns_host/
command_host.rs

1//! Store-backed command context and completion implementation.
2
3use std::collections::{HashMap, HashSet};
4use std::sync::Arc;
5
6use async_trait::async_trait;
7use uuid::Uuid;
8
9use everruns_core::capabilities::CapabilityRegistry;
10use everruns_core::command_host::{
11    CommandHost, CommandTurnContext, SessionCompletion, SessionCompletionError,
12    SessionCompletionRequest, SessionCompletionStream,
13};
14use everruns_core::execution_loading::{AgentStore, HarnessStore, SessionStore};
15use everruns_core::file_services::{FileResolver, ResolvedFile};
16use everruns_core::image_services::{ImageResolver, ResolvedImage};
17use everruns_core::message::{Controls, Message, MessageRole, patch_dangling_tool_calls};
18use everruns_core::message_retriever::MessageRetriever;
19use everruns_core::provider_resolution::ProviderStore;
20use everruns_core::runtime_context::{AssembledTurnContext, ResolvedModelExecution};
21use everruns_core::session_files::SessionFileSystem;
22use everruns_provider::driver_registry::{
23    ChatDriver, DriverRegistry, LlmCallConfig, LlmMessage, LlmMessageRole, ToolSearchConfig,
24};
25use everruns_provider::error::{AgentLoopError, Result};
26use everruns_provider::runtime_provider::ProviderEndpoint;
27use everruns_provider::typed_id::SessionId;
28use everruns_provider::user_facing_error::UserFacingErrorContext;
29
30use crate::runtime_context::{inspect_turn_context_for_session, resolve_model_execution};
31
32/// Store-backed [`CommandHost`] shared by embedded and durable hosts.
33///
34/// One instance is built per command dispatch. Its assembled context is
35/// memoized so `turn_context()` and completion share one store resolution.
36pub struct StoreCommandHost {
37    session_id: SessionId,
38    harness_store: Arc<dyn HarnessStore>,
39    agent_store: Arc<dyn AgentStore>,
40    session_store: Arc<dyn SessionStore>,
41    message_retriever: Arc<dyn MessageRetriever>,
42    provider_store: Arc<dyn ProviderStore>,
43    capability_registry: CapabilityRegistry,
44    driver_registry: DriverRegistry,
45    image_resolver: Option<Arc<dyn ImageResolver>>,
46    file_resolver: Option<Arc<dyn FileResolver>>,
47    file_store: Option<Arc<dyn SessionFileSystem>>,
48    assembled: tokio::sync::OnceCell<AssembledTurnContext>,
49}
50
51impl StoreCommandHost {
52    /// Construct a command host from org-scoped runtime stores and registries.
53    #[allow(clippy::too_many_arguments)]
54    pub fn new(
55        session_id: SessionId,
56        harness_store: Arc<dyn HarnessStore>,
57        agent_store: Arc<dyn AgentStore>,
58        session_store: Arc<dyn SessionStore>,
59        message_retriever: Arc<dyn MessageRetriever>,
60        provider_store: Arc<dyn ProviderStore>,
61        capability_registry: CapabilityRegistry,
62        driver_registry: DriverRegistry,
63    ) -> Self {
64        Self {
65            session_id,
66            harness_store,
67            agent_store,
68            session_store,
69            message_retriever,
70            provider_store,
71            capability_registry,
72            driver_registry,
73            image_resolver: None,
74            file_resolver: None,
75            file_store: None,
76            assembled: tokio::sync::OnceCell::new(),
77        }
78    }
79
80    /// Resolve `image_file` references for provider conversion.
81    pub fn with_image_resolver(mut self, image_resolver: Arc<dyn ImageResolver>) -> Self {
82        self.image_resolver = Some(image_resolver);
83        self
84    }
85
86    pub fn with_file_resolver(mut self, file_resolver: Arc<dyn FileResolver>) -> Self {
87        self.file_resolver = Some(file_resolver);
88        self
89    }
90
91    /// Supply the session filesystem used by dynamic prompt capabilities.
92    pub fn with_file_store(mut self, file_store: Arc<dyn SessionFileSystem>) -> Self {
93        self.file_store = Some(file_store);
94        self
95    }
96
97    /// Seed a context already assembled by the dispatching host.
98    pub fn with_assembled_context(mut self, assembled: AssembledTurnContext) -> Self {
99        self.assembled = tokio::sync::OnceCell::new_with(Some(assembled));
100        self
101    }
102
103    async fn assembled(&self) -> Result<&AssembledTurnContext> {
104        self.assembled
105            .get_or_try_init(|| async {
106                inspect_turn_context_for_session(
107                    self.harness_store.as_ref(),
108                    self.agent_store.as_ref(),
109                    self.session_store.as_ref(),
110                    self.message_retriever.as_ref(),
111                    self.provider_store.as_ref(),
112                    &self.capability_registry,
113                    &self.driver_registry,
114                    self.session_id,
115                    &[],
116                    self.file_store.clone(),
117                )
118                .await
119            })
120            .await
121    }
122
123    async fn resolve_images(&self, messages: &[Message]) -> HashMap<Uuid, ResolvedImage> {
124        let Some(resolver) = &self.image_resolver else {
125            return HashMap::new();
126        };
127        let image_ids: HashSet<Uuid> = messages
128            .iter()
129            .flat_map(everruns_core::llm_conversions::extract_image_file_ids)
130            .collect();
131        let mut resolved = HashMap::new();
132        for image_id in image_ids {
133            if let Ok(Some(image)) = resolver.resolve_image(image_id).await {
134                resolved.insert(image_id, image);
135            }
136        }
137        resolved
138    }
139
140    async fn resolve_files(&self, messages: &[Message]) -> HashMap<Uuid, ResolvedFile> {
141        let Some(resolver) = &self.file_resolver else {
142            return HashMap::new();
143        };
144        let file_ids: HashSet<Uuid> = messages
145            .iter()
146            .flat_map(everruns_core::llm_conversions::extract_file_ids)
147            .collect();
148        match resolver
149            .resolve_files(&file_ids.into_iter().collect::<Vec<_>>())
150            .await
151        {
152            Ok(map) => map,
153            Err(e) => {
154                tracing::warn!("Failed to resolve file attachments: {e}");
155                HashMap::new()
156            }
157        }
158    }
159
160    async fn resolve_completion_model(
161        &self,
162        controls: Option<&Controls>,
163        assembled: &AssembledTurnContext,
164    ) -> std::result::Result<ResolvedModelExecution, SessionCompletionError> {
165        let requested = controls.and_then(|controls| controls.model_id);
166        if requested.is_none() || requested == assembled.resolved_model_id {
167            return Ok(assembled.model.clone());
168        }
169        let model_id = requested.expect("checked above");
170        let spec = self
171            .provider_store
172            .get_model_spec(model_id)
173            .await
174            .map_err(SessionCompletionError::InvalidRequest)?
175            .ok_or_else(|| {
176                SessionCompletionError::InvalidRequest(AgentLoopError::config(format!(
177                    "Model not found: {model_id}"
178                )))
179            })?;
180        let error_context = UserFacingErrorContext::default()
181            .with_provider(spec.provider.to_string())
182            .with_model_id(spec.model.clone());
183        resolve_model_execution(self.provider_store.as_ref(), &self.driver_registry, spec)
184            .await
185            .map_err(|error| SessionCompletionError::Completion {
186                error: error.to_string(),
187                context: error_context,
188            })
189    }
190
191    async fn prepare_completion(
192        &self,
193        request: SessionCompletionRequest,
194    ) -> std::result::Result<PreparedCompletion, SessionCompletionError> {
195        let assembled = self
196            .assembled()
197            .await
198            .map_err(SessionCompletionError::InvalidRequest)?;
199        let model = self
200            .resolve_completion_model(request.controls.as_ref(), assembled)
201            .await?;
202        let context = UserFacingErrorContext::default()
203            .with_provider(model.provider_type.to_string())
204            .with_model_id(model.model.clone());
205        let messages = patch_dangling_tool_calls(&request.messages);
206        let resolved_images = self.resolve_images(&messages).await;
207        let resolved_files = self.resolve_files(&messages).await;
208        let mut llm_messages: Vec<LlmMessage> = request
209            .system_prompts
210            .iter()
211            .filter(|prompt| !prompt.is_empty())
212            .map(|prompt| LlmMessage::text(LlmMessageRole::System, prompt.clone()))
213            .collect();
214        for message in &messages {
215            let mut llm_message =
216                everruns_core::llm_conversions::llm_message_from_message_with_attachments(
217                    message,
218                    &resolved_images,
219                    &resolved_files,
220                );
221            if message.role == MessageRole::User
222                && let Some(actor) = &message.external_actor
223            {
224                llm_message.prepend_text_prefix(&format!("[{}] ", actor.display_label()));
225            }
226            llm_messages.push(llm_message);
227        }
228        let mut builder = everruns_core::llm_conversions::llm_call_config_builder_from_agent(
229            &assembled.runtime_agent,
230        )
231        .model(&model.model)
232        .tools(vec![])
233        .tool_search(ToolSearchConfig {
234            enabled: false,
235            threshold: usize::MAX,
236        })
237        .previous_response_id(None)
238        .with_metadata("session_id", self.session_id.to_string());
239        if let Some(effort) = request
240            .controls
241            .as_ref()
242            .and_then(|controls| controls.reasoning.as_ref())
243            .and_then(|reasoning| reasoning.effort)
244        {
245            builder = builder.reasoning_effort(effort);
246        }
247        for (key, value) in &request.metadata {
248            builder = builder.with_metadata(key, value);
249        }
250        Ok(PreparedCompletion {
251            llm_messages,
252            llm_config: builder.build(),
253            driver: model.driver,
254            context,
255        })
256    }
257}
258
259struct PreparedCompletion {
260    llm_messages: Vec<LlmMessage>,
261    llm_config: LlmCallConfig,
262    driver: Arc<dyn ChatDriver>,
263    context: UserFacingErrorContext,
264}
265
266#[async_trait]
267impl CommandHost for StoreCommandHost {
268    async fn turn_context(&self) -> Result<CommandTurnContext> {
269        let assembled = self.assembled().await?;
270        Ok(CommandTurnContext {
271            session_id: assembled.snapshot.session_id,
272            messages: assembled.messages.clone(),
273            system_prompt: assembled.runtime_agent.system_prompt.clone(),
274            model: assembled.model.model.clone(),
275            provider_type: assembled.model.provider_type.to_string(),
276            resolved_locale: assembled.resolved_locale.clone(),
277        })
278    }
279
280    async fn completion(
281        &self,
282        request: SessionCompletionRequest,
283    ) -> std::result::Result<SessionCompletion, SessionCompletionError> {
284        let prepared = self.prepare_completion(request).await?;
285        let completion_error = |error: String| SessionCompletionError::Completion {
286            error,
287            context: prepared.context.clone(),
288        };
289        let response = prepared
290            .driver
291            .chat_completion(
292                &ProviderEndpoint::default(),
293                prepared.llm_messages,
294                &prepared.llm_config,
295            )
296            .await
297            .map_err(|error| completion_error(error.to_string()))?;
298        let text = response.text.trim().to_string();
299        if text.is_empty() {
300            return Err(completion_error(
301                "session completion returned an empty response".to_string(),
302            ));
303        }
304        Ok(SessionCompletion { text })
305    }
306
307    async fn completion_stream(
308        &self,
309        request: SessionCompletionRequest,
310    ) -> std::result::Result<SessionCompletionStream, SessionCompletionError> {
311        let prepared = self.prepare_completion(request).await?;
312        let events = prepared
313            .driver
314            .chat_completion_stream(
315                &ProviderEndpoint::default(),
316                prepared.llm_messages,
317                &prepared.llm_config,
318            )
319            .await
320            .map_err(|error| SessionCompletionError::Completion {
321                error: error.to_string(),
322                context: prepared.context.clone(),
323            })?;
324        Ok(SessionCompletionStream {
325            events,
326            context: prepared.context,
327        })
328    }
329}