1use 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
32pub 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 #[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 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 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 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}