1use crate::llm::transport::LlmTransportError;
2use crate::llm::types::{
3 AttachmentSource, LlmContentBlock, LlmEventSender, LlmJsonSchema, LlmMessage, LlmOutputSpec,
4 LlmRequest, LlmRequestScope, LlmResponse, LlmRole, LlmStreamEvent, LlmTerminalReason,
5 LlmToolChoice,
6};
7use crate::provider::{ModelCapability, ModelEffortValidationCategory, ProviderHandle};
8use crate::{LashSchema, SchemaContract};
9use lash_trace::{TraceContext, TraceError, TraceEvent, TraceSink};
10use std::sync::Arc;
11
12#[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
13#[serde(rename_all = "snake_case")]
14pub enum DirectRole {
15 System,
16 User,
17 Assistant,
18}
19
20#[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
21pub enum DirectPart {
22 Text(String),
23 Attachment(usize),
24}
25
26#[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
27pub struct DirectMessage {
28 pub role: DirectRole,
29 pub parts: Vec<DirectPart>,
30}
31
32#[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
33pub struct DirectJsonSchema {
34 pub name: String,
35 pub schema: SchemaContract,
36 pub strict: bool,
37}
38
39#[derive(Clone, Debug, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
40pub enum DirectOutputSpec {
41 #[default]
42 Text,
43 JsonObject,
44 JsonSchema(DirectJsonSchema),
45}
46
47#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
48pub struct DirectRequest {
49 pub model: String,
50 #[serde(default)]
51 pub model_variant: crate::ReasoningSelection,
52 #[serde(default, skip_serializing_if = "ModelCapability::is_empty")]
53 pub model_capability: ModelCapability,
54 #[serde(default)]
55 pub messages: Vec<DirectMessage>,
56 #[serde(default)]
57 pub attachments: Vec<AttachmentSource>,
58 #[serde(default)]
59 pub output: DirectOutputSpec,
60 #[serde(default)]
61 pub generation: crate::GenerationOptions,
62 #[serde(default, skip)]
63 pub stream_events: Option<LlmEventSender>,
64 #[serde(default, skip_serializing_if = "Option::is_none")]
65 pub session_id: Option<String>,
66 #[serde(default, skip_serializing_if = "Option::is_none")]
67 pub caused_by: Option<crate::CausalRef>,
68 #[serde(default, skip_serializing_if = "Option::is_none")]
69 pub replay: Option<crate::RuntimeReplay>,
70}
71
72impl DirectRequest {
73 pub fn text(model: impl Into<String>, prompt: impl Into<String>) -> Self {
74 Self {
75 model: model.into(),
76 model_variant: crate::ReasoningSelection::ProviderDefault,
77 model_capability: ModelCapability::default(),
78 messages: vec![DirectMessage {
79 role: DirectRole::User,
80 parts: vec![DirectPart::Text(prompt.into())],
81 }],
82 attachments: Vec::new(),
83 output: DirectOutputSpec::Text,
84 generation: crate::GenerationOptions::default(),
85 stream_events: None,
86 session_id: None,
87 caused_by: None,
88 replay: None,
89 }
90 }
91
92 pub fn json(model: impl Into<String>, prompt: impl Into<String>) -> Self {
93 Self {
94 output: DirectOutputSpec::JsonObject,
95 ..Self::text(model, prompt)
96 }
97 }
98
99 pub fn json_schema(
100 model: impl Into<String>,
101 prompt: impl Into<String>,
102 schema: DirectJsonSchema,
103 ) -> Self {
104 Self {
105 output: DirectOutputSpec::JsonSchema(schema),
106 ..Self::text(model, prompt)
107 }
108 }
109
110 pub fn with_replay_key(mut self, key: impl Into<String>) -> Self {
111 self.replay = Some(crate::RuntimeReplay { key: key.into() });
112 self
113 }
114
115 pub fn with_caused_by(mut self, caused_by: crate::CausalRef) -> Self {
116 self.caused_by = Some(caused_by);
117 self
118 }
119}
120
121#[derive(Debug, thiserror::Error, Clone)]
122pub enum DirectLlmError {
123 #[error("invalid request: {message}")]
124 InvalidRequest {
125 category: ModelEffortValidationCategory,
126 message: String,
127 },
128 #[error("invalid response: {0}")]
129 InvalidResponse(String),
130 #[error("transport error: {0}")]
131 Transport(#[from] Box<LlmTransportError>),
132}
133
134#[derive(Clone, Debug)]
137pub struct DirectLlmResult {
138 pub response: LlmResponse,
139 pub llm_call: crate::LlmCallRecord,
140}
141
142impl std::ops::Deref for DirectLlmResult {
143 type Target = LlmResponse;
144
145 fn deref(&self) -> &Self::Target {
146 &self.response
147 }
148}
149
150impl DirectLlmResult {
151 pub fn into_response(self) -> LlmResponse {
152 self.response
153 }
154}
155
156pub struct DirectLlmClient {
157 provider: ProviderHandle,
158 trace_sink: Option<Arc<dyn TraceSink>>,
159 trace_context: TraceContext,
160 clock: Arc<dyn crate::Clock>,
161}
162
163impl DirectLlmClient {
164 pub fn new(provider: ProviderHandle) -> Self {
165 Self {
166 provider,
167 trace_sink: None,
168 trace_context: TraceContext::default(),
169 clock: Arc::new(crate::SystemClock),
170 }
171 }
172
173 pub fn with_trace_sink(mut self, sink: Option<Arc<dyn TraceSink>>) -> Self {
174 self.trace_sink = sink;
175 self
176 }
177
178 pub fn with_trace_context(mut self, context: TraceContext) -> Self {
179 self.trace_context = context;
180 self
181 }
182
183 pub fn with_clock(mut self, clock: Arc<dyn crate::Clock>) -> Self {
184 self.clock = clock;
185 self
186 }
187
188 pub fn provider(&self) -> &ProviderHandle {
189 &self.provider
190 }
191
192 pub fn provider_mut(&mut self) -> &mut ProviderHandle {
193 &mut self.provider
194 }
195
196 pub async fn complete(
197 &mut self,
198 mut request: DirectRequest,
199 ) -> Result<DirectLlmResult, DirectLlmError> {
200 request.model_variant = request
204 .model_capability
205 .validate_selection(&request.model, self.provider.kind(), &request.model_variant)
206 .map_err(|error| DirectLlmError::InvalidRequest {
207 category: error.category,
208 message: error.message,
209 })?;
210
211 let output_for_validation = request.output.clone();
212 let model = request.model.clone();
213 let llm_request = build_llm_request(&self.provider, request, model);
214 let llm_call_id = if self.trace_sink.is_some() {
215 let id = uuid::Uuid::new_v4().to_string();
216 crate::trace::emit_trace(
217 &self.trace_sink,
218 &self.trace_context,
219 TraceContext::default().for_llm_call(id.clone()),
220 TraceEvent::LlmCallStarted {
221 request: crate::trace::trace_llm_request(&llm_request),
222 },
223 self.clock.as_ref(),
224 );
225 Some(id)
226 } else {
227 None
228 };
229 match self.provider.complete(llm_request).await {
230 Ok(response) => {
231 if let Err(error) = validate_direct_output(&output_for_validation, &response) {
232 if let Some(llm_call_id) = llm_call_id {
233 crate::trace::emit_trace(
234 &self.trace_sink,
235 &self.trace_context,
236 TraceContext::default().for_llm_call(llm_call_id),
237 TraceEvent::LlmCallFailed {
238 error: TraceError {
239 message: error.to_string(),
240 retryable: false,
241 terminal_reason: Some(
242 LlmTerminalReason::ProviderError.code().to_string(),
243 ),
244 code: Some("invalid_structured_output".to_string()),
245 raw: None,
246 },
247 stream_summary: None,
248 },
249 self.clock.as_ref(),
250 );
251 }
252 return Err(error);
253 }
254 if let Some(llm_call_id) = llm_call_id {
255 crate::trace::emit_trace(
256 &self.trace_sink,
257 &self.trace_context,
258 TraceContext::default().for_llm_call(llm_call_id),
259 TraceEvent::LlmCallCompleted {
260 response: crate::trace::trace_llm_response(
261 response.full_text.clone(),
262 0,
263 Some(response.terminal_reason),
264 crate::trace::trace_output_parts(&response.parts),
265 ),
266 usage: Some(crate::trace::trace_usage_from_llm(&response.usage)),
267 provider_usage: response.provider_usage.clone(),
268 stream_summary: None,
269 },
270 self.clock.as_ref(),
271 );
272 }
273 Ok(DirectLlmResult {
274 response: response.response,
275 llm_call: response.call_record,
276 })
277 }
278 Err(error) => {
279 if let Some(llm_call_id) = llm_call_id {
280 crate::trace::emit_trace(
281 &self.trace_sink,
282 &self.trace_context,
283 TraceContext::default().for_llm_call(llm_call_id),
284 TraceEvent::LlmCallFailed {
285 error: TraceError {
286 message: error.message.clone(),
287 retryable: error.retryable,
288 terminal_reason: Some(error.terminal_reason.code().to_string()),
289 code: error.code.clone(),
290 raw: error.raw.as_deref().cloned(),
291 },
292 stream_summary: None,
293 },
294 self.clock.as_ref(),
295 );
296 }
297 Err(DirectLlmError::from(Box::new(error.error)))
298 }
299 }
300 }
301}
302
303pub(crate) fn build_llm_request(
304 provider: &ProviderHandle,
305 request: DirectRequest,
306 model: String,
307) -> LlmRequest {
308 let stream_events = transport_stream_events_for_direct(provider, request.stream_events);
309 let DirectRequest {
310 model: _,
311 model_variant,
312 model_capability,
313 messages,
314 attachments,
315 output,
316 generation,
317 stream_events: _,
318 session_id,
319 caused_by: _,
320 replay: _,
321 } = request;
322
323 let output_spec = match output {
324 DirectOutputSpec::Text => None,
325 DirectOutputSpec::JsonObject => Some(LlmOutputSpec::JsonObject),
326 DirectOutputSpec::JsonSchema(schema) => Some(LlmOutputSpec::JsonSchema(LlmJsonSchema {
327 name: schema.name,
328 schema: schema.schema,
329 strict: schema.strict,
330 })),
331 };
332
333 let mut llm_messages = Vec::new();
334 for message in messages {
335 let role = match message.role {
336 DirectRole::System => LlmRole::System,
337 DirectRole::User => LlmRole::User,
338 DirectRole::Assistant => LlmRole::Assistant,
339 };
340 let mut blocks: Vec<LlmContentBlock> = Vec::new();
341 for part in message.parts {
342 match part {
343 DirectPart::Text(text) => {
344 if !text.is_empty() {
345 blocks.push(LlmContentBlock::Text {
346 text: text.into(),
347 response_meta: None,
348 cache_breakpoint: false,
349 });
350 }
351 }
352 DirectPart::Attachment(idx) => {
353 blocks.push(LlmContentBlock::Attachment {
354 attachment_idx: idx,
355 });
356 }
357 }
358 }
359 if !blocks.is_empty() {
360 llm_messages.push(LlmMessage::new(role, blocks));
361 }
362 }
363
364 let scope = match session_id {
365 Some(session_id) => LlmRequestScope::new(
366 session_id.clone(),
367 format!("{session_id}:frame:direct"),
368 format!("{session_id}:direct"),
369 ),
370 None => {
371 let request_id = uuid::Uuid::new_v4().to_string();
372 LlmRequestScope::new(
373 format!("direct:{request_id}"),
374 format!("direct:{request_id}:frame"),
375 request_id,
376 )
377 }
378 };
379
380 LlmRequest {
381 model,
382 messages: llm_messages,
383 attachments,
384 resolved_stored: Default::default(),
385 tools: Vec::new().into(),
386 tool_choice: LlmToolChoice::None,
387 model_variant,
388 model_capability,
389 generation,
390 scope,
391 output_spec,
392 stream_events,
393 provider_trace: None,
394 }
395}
396
397fn validate_direct_output(
398 output: &DirectOutputSpec,
399 response: &LlmResponse,
400) -> Result<(), DirectLlmError> {
401 let DirectOutputSpec::JsonSchema(schema) = output else {
402 return Ok(());
403 };
404 let parsed: serde_json::Value = serde_json::from_str(response.full_text.trim())
405 .map_err(|err| DirectLlmError::InvalidResponse(format!("expected JSON: {err}")))?;
406 LashSchema::new(schema.schema.canonical().clone())
407 .validate(&parsed)
408 .map_err(DirectLlmError::InvalidResponse)
409}
410
411fn transport_stream_events_for_direct(
412 provider: &ProviderHandle,
413 requested: Option<LlmEventSender>,
414) -> Option<LlmEventSender> {
415 if requested.is_some() {
416 return requested;
417 }
418 if provider.requires_streaming() {
419 Some(LlmEventSender::new(|_event: LlmStreamEvent| {}))
420 } else {
421 None
422 }
423}
424
425#[cfg(test)]
426mod tests {
427 use super::*;
428 use crate::llm::types::{LlmOutputPart, LlmTerminalReason, LlmUsage};
429 use crate::provider::{ProviderOptions, ProviderReliability};
430 use crate::testing::TestProvider;
431 use serde_json::json;
432 use std::sync::{Arc, Mutex};
433
434 #[test]
435 fn json_schema_request_preserves_output_schema() {
436 let schema = DirectJsonSchema {
437 name: "answer_shape".to_string(),
438 schema: json!({
439 "type": "object",
440 "properties": {
441 "answer": { "type": "string" }
442 },
443 "required": ["answer"]
444 })
445 .into(),
446 strict: true,
447 };
448
449 let request = DirectRequest::json_schema("model-a", "return json", schema.clone());
450
451 assert_eq!(
452 request.output,
453 DirectOutputSpec::JsonSchema(schema),
454 "DirectRequest::json_schema must carry the requested output schema"
455 );
456 }
457
458 #[test]
459 fn direct_client_provider_accessors_expose_owned_provider_handle() {
460 let provider = TestProvider::builder()
461 .kind("direct-accessor-provider")
462 .serialize_config(|| json!({"provider": "owned"}))
463 .build()
464 .into_handle();
465 let mut client = DirectLlmClient::new(provider);
466
467 assert_eq!(client.provider().kind(), "direct-accessor-provider");
468 assert_eq!(
469 client.provider().to_spec().config,
470 json!({"provider": "owned"})
471 );
472
473 let options = ProviderOptions {
474 reliability: ProviderReliability::default().max_attempts(7),
475 max_output_tokens: Some(123),
476 ..Default::default()
477 };
478 client.provider_mut().set_options(options.clone());
479
480 assert_eq!(client.provider().options(), options);
481 }
482
483 #[tokio::test]
484 async fn direct_client_complete_delegates_to_provider_and_returns_response() {
485 let captured_request: Arc<Mutex<Option<LlmRequest>>> = Arc::new(Mutex::new(None));
486 let captured_for_provider = Arc::clone(&captured_request);
487 let provider = TestProvider::builder()
488 .kind("direct-complete-provider")
489 .complete(move |request| {
490 let captured_for_provider = Arc::clone(&captured_for_provider);
491 async move {
492 *captured_for_provider.lock().expect("capture lock") = Some(request);
493 Ok(LlmResponse {
494 full_text: "provider delegated response".to_string(),
495 parts: vec![LlmOutputPart::Text {
496 text: "provider delegated response".to_string(),
497 response_meta: None,
498 }],
499 usage: LlmUsage {
500 input_tokens: 11,
501 output_tokens: 3,
502 ..Default::default()
503 },
504 terminal_reason: LlmTerminalReason::Stop,
505 response_metadata: Default::default(),
506 ..Default::default()
507 })
508 }
509 })
510 .build()
511 .into_handle();
512 let mut client = DirectLlmClient::new(provider);
513 let mut request = DirectRequest::json("direct-model", "answer as json");
514 request.session_id = Some("direct-session".to_string());
515
516 let response = client
517 .complete(request)
518 .await
519 .expect("direct completion should delegate");
520
521 assert_eq!(response.full_text, "provider delegated response");
522 assert_eq!(response.llm_call.attempts.len(), 1);
523 let captured = captured_request
524 .lock()
525 .expect("capture lock")
526 .clone()
527 .expect("provider should receive a request");
528 assert_eq!(captured.model, "direct-model");
529 assert_eq!(captured.scope.session_id, "direct-session");
530 assert_eq!(captured.scope.agent_frame_id, "direct-session:frame:direct");
531 assert_eq!(captured.scope.request_id, "direct-session:direct");
532 assert!(matches!(
533 captured.output_spec,
534 Some(LlmOutputSpec::JsonObject)
535 ));
536 assert_eq!(captured.messages.len(), 1);
537 }
538
539 #[tokio::test]
540 async fn direct_client_validates_json_schema_output_against_canonical_schema() {
541 let provider = TestProvider::builder()
542 .kind("direct-validation-provider")
543 .complete(|_request| async {
544 Ok(LlmResponse {
545 full_text: r#"{"items":[]}"#.to_string(),
546 terminal_reason: LlmTerminalReason::Stop,
547 response_metadata: Default::default(),
548 ..Default::default()
549 })
550 })
551 .build()
552 .into_handle();
553 let mut client = DirectLlmClient::new(provider);
554 let request = DirectRequest::json_schema(
555 "direct-model",
556 "return items",
557 DirectJsonSchema {
558 name: "items_result".to_string(),
559 schema: json!({
560 "type": "object",
561 "required": ["items"],
562 "properties": {
563 "items": {
564 "type": "array",
565 "minItems": 1,
566 "items": { "type": "string" }
567 }
568 }
569 })
570 .into(),
571 strict: true,
572 },
573 );
574
575 let err = client
576 .complete(request)
577 .await
578 .expect_err("empty items must fail canonical validation");
579
580 assert!(matches!(err, DirectLlmError::InvalidResponse(_)));
581 let error = err.to_string();
582 assert!(
583 error.contains("items") && error.contains("[] has less than 1 item"),
584 "{error}"
585 );
586 }
587
588 fn reasoning_capability() -> ModelCapability {
589 ModelCapability {
590 reasoning: Some(crate::ReasoningCapability {
591 efforts: ["low", "medium", "high", "max"]
592 .into_iter()
593 .map(String::from)
594 .collect(),
595 aliases: std::collections::BTreeMap::from([(
596 "xhigh".to_string(),
597 "max".to_string(),
598 )]),
599 ..Default::default()
600 }),
601 cache_control: None,
602 stream_termination: None,
603 }
604 }
605
606 #[tokio::test]
607 async fn direct_client_rejects_unsupported_effort_before_provider_call() {
608 let called = Arc::new(Mutex::new(false));
609 let called_for_provider = Arc::clone(&called);
610 let provider = TestProvider::builder()
611 .kind("direct-reject")
612 .complete(move |_request| {
613 let called = Arc::clone(&called_for_provider);
614 async move {
615 *called.lock().expect("called lock") = true;
616 Ok(LlmResponse::default())
617 }
618 })
619 .build()
620 .into_handle();
621 let mut client = DirectLlmClient::new(provider);
622
623 let mut request = DirectRequest::text("direct-model", "hi");
624 request.model_variant = crate::ReasoningSelection::Effort("turbo".to_string());
625 request.model_capability = reasoning_capability();
626
627 let err = client
628 .complete(request)
629 .await
630 .expect_err("unsupported effort must be rejected");
631 assert!(matches!(
632 err,
633 DirectLlmError::InvalidRequest {
634 category: ModelEffortValidationCategory::UnsupportedEffort,
635 ..
636 }
637 ));
638 assert!(err.to_string().contains("Unsupported effort `turbo`"));
639 assert!(
640 !*called.lock().expect("called lock"),
641 "the provider must not be called when the effort is rejected"
642 );
643 }
644
645 #[tokio::test]
646 async fn direct_client_normalizes_alias_effort_into_outgoing_request() {
647 let captured: Arc<Mutex<Option<crate::ReasoningSelection>>> = Arc::new(Mutex::new(None));
648 let captured_for_provider = Arc::clone(&captured);
649 let provider = TestProvider::builder()
650 .kind("direct-alias")
651 .complete(move |request| {
652 let captured = Arc::clone(&captured_for_provider);
653 async move {
654 *captured.lock().expect("capture lock") = Some(request.model_variant.clone());
655 Ok(LlmResponse {
656 full_text: "ok".to_string(),
657 terminal_reason: LlmTerminalReason::Stop,
658 response_metadata: Default::default(),
659 ..Default::default()
660 })
661 }
662 })
663 .build()
664 .into_handle();
665 let mut client = DirectLlmClient::new(provider);
666
667 let mut request = DirectRequest::text("direct-model", "hi");
668 request.model_variant = crate::ReasoningSelection::Effort("XHigh".to_string());
669 request.model_capability = reasoning_capability();
670
671 client.complete(request).await.expect("completion");
672 let seen = captured
673 .lock()
674 .expect("capture lock")
675 .clone()
676 .expect("provider must be called");
677 assert_eq!(
678 seen,
679 crate::ReasoningSelection::Effort("max".to_string()),
680 "alias `XHigh` must clamp to canonical `max` before the provider sees the request"
681 );
682 }
683
684 #[tokio::test]
685 async fn direct_client_rejects_effort_when_model_is_not_configurable() {
686 let provider = TestProvider::builder()
687 .kind("direct-not-configurable")
688 .complete(|_request| async { Ok(LlmResponse::default()) })
689 .build()
690 .into_handle();
691 let mut client = DirectLlmClient::new(provider);
692
693 let mut request = DirectRequest::text("direct-model", "hi");
694 request.model_variant = crate::ReasoningSelection::Effort("high".to_string());
695 let err = client
698 .complete(request)
699 .await
700 .expect_err("effort on a non-configurable model must be rejected");
701 assert!(matches!(
702 err,
703 DirectLlmError::InvalidRequest {
704 category: ModelEffortValidationCategory::EffortNotConfigurable,
705 ..
706 }
707 ));
708 }
709
710 #[tokio::test]
711 async fn direct_client_rejects_missing_mandatory_effort() {
712 let provider = TestProvider::builder()
713 .kind("direct-mandatory")
714 .complete(|_request| async { Ok(LlmResponse::default()) })
715 .build()
716 .into_handle();
717 let mut client = DirectLlmClient::new(provider);
718
719 let mut capability = reasoning_capability();
720 capability.reasoning.as_mut().expect("reasoning").mandatory = true;
721 let mut request = DirectRequest::text("direct-model", "hi");
722 request.model_capability = capability;
723 let err = client
726 .complete(request)
727 .await
728 .expect_err("missing mandatory effort must be rejected");
729 assert!(matches!(
730 err,
731 DirectLlmError::InvalidRequest {
732 category: ModelEffortValidationCategory::EffortRequired,
733 ..
734 }
735 ));
736 }
737
738 #[test]
739 fn build_llm_request_preserves_nonempty_content_and_drops_empty_messages() {
740 let provider = TestProvider::default().into_handle();
741 let request = DirectRequest {
742 model: "input-model".to_string(),
743 messages: vec![
744 DirectMessage {
745 role: DirectRole::System,
746 parts: vec![DirectPart::Text(String::new())],
747 },
748 DirectMessage {
749 role: DirectRole::User,
750 parts: vec![
751 DirectPart::Text("hello".to_string()),
752 DirectPart::Text(String::new()),
753 ],
754 },
755 DirectMessage {
756 role: DirectRole::Assistant,
757 parts: vec![DirectPart::Attachment(2)],
758 },
759 ],
760 attachments: Vec::new(),
761 output: DirectOutputSpec::Text,
762 generation: crate::GenerationOptions::default(),
763 stream_events: None,
764 session_id: None,
765 model_variant: Default::default(),
766 model_capability: ModelCapability::default(),
767 caused_by: None,
768 replay: None,
769 };
770
771 let llm_request = build_llm_request(&provider, request, "transport-model".to_string());
772
773 assert_eq!(llm_request.model, "transport-model");
774 assert_eq!(
775 llm_request.messages.len(),
776 2,
777 "empty normalized messages must be dropped"
778 );
779 assert_eq!(llm_request.messages[0].role, LlmRole::User);
780 assert_eq!(llm_request.messages[0].blocks.len(), 1);
781 assert!(matches!(
782 &llm_request.messages[0].blocks[0],
783 LlmContentBlock::Text { text, .. } if text.as_ref() == "hello"
784 ));
785 assert_eq!(llm_request.messages[1].role, LlmRole::Assistant);
786 assert!(matches!(
787 &llm_request.messages[1].blocks[0],
788 LlmContentBlock::Attachment { attachment_idx: 2 }
789 ));
790 }
791
792 #[test]
793 fn build_llm_request_preserves_direct_stream_sender_and_adds_required_noop_sender() {
794 let captured_events: Arc<Mutex<Vec<LlmStreamEvent>>> = Arc::new(Mutex::new(Vec::new()));
795 let captured_for_sender = Arc::clone(&captured_events);
796 let requested_sender = LlmEventSender::new(move |event| {
797 captured_for_sender
798 .lock()
799 .expect("stream event lock")
800 .push(event);
801 });
802 let mut request = DirectRequest::text("model", "prompt");
803 request.stream_events = Some(requested_sender);
804 let provider = TestProvider::default().into_handle();
805
806 let llm_request = build_llm_request(&provider, request, "model".to_string());
807 let sender = llm_request
808 .stream_events
809 .expect("explicit direct stream sender must be preserved");
810 sender.send(LlmStreamEvent::Delta("delta".to_string()));
811 assert_eq!(captured_events.lock().expect("stream event lock").len(), 1);
812
813 let streaming_provider = TestProvider::builder()
814 .requires_streaming(true)
815 .build()
816 .into_handle();
817 let llm_request = build_llm_request(
818 &streaming_provider,
819 DirectRequest::text("model", "prompt"),
820 "model".to_string(),
821 );
822 assert!(
823 llm_request.stream_events.is_some(),
824 "providers that require streaming need a no-op sender even when direct caller did not request one"
825 );
826 }
827}