Skip to main content

lash_core/
direct.rs

1use crate::llm::transport::LlmTransportError;
2use crate::llm::types::{
3    LlmAttachment, 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    Image(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<LlmAttachment>,
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/// Successful single-shot direct LLM result with the sealed provider-attempt
135/// history that produced it.
136#[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        // Validate the requested effort against the capability that travels
201        // with the request, and write the resolved (alias-normalized) effort
202        // back so the provider never sees an un-clamped value.
203        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.clone(),
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::Image(idx) => {
353                    blocks.push(LlmContentBlock::Image {
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        tools: Vec::new().into(),
385        tool_choice: LlmToolChoice::None,
386        model_variant,
387        model_capability,
388        generation,
389        scope,
390        output_spec,
391        stream_events,
392        provider_trace: None,
393    }
394}
395
396fn validate_direct_output(
397    output: &DirectOutputSpec,
398    response: &LlmResponse,
399) -> Result<(), DirectLlmError> {
400    let DirectOutputSpec::JsonSchema(schema) = output else {
401        return Ok(());
402    };
403    let parsed: serde_json::Value = serde_json::from_str(response.full_text.trim())
404        .map_err(|err| DirectLlmError::InvalidResponse(format!("expected JSON: {err}")))?;
405    LashSchema::new(schema.schema.canonical().clone())
406        .validate(&parsed)
407        .map_err(DirectLlmError::InvalidResponse)
408}
409
410fn transport_stream_events_for_direct(
411    provider: &ProviderHandle,
412    requested: Option<LlmEventSender>,
413) -> Option<LlmEventSender> {
414    if requested.is_some() {
415        return requested;
416    }
417    if provider.requires_streaming() {
418        Some(LlmEventSender::new(|_event: LlmStreamEvent| {}))
419    } else {
420        None
421    }
422}
423
424#[cfg(test)]
425mod tests {
426    use super::*;
427    use crate::llm::types::{LlmOutputPart, LlmTerminalReason, LlmUsage};
428    use crate::provider::{ProviderOptions, ProviderReliability};
429    use crate::testing::TestProvider;
430    use serde_json::json;
431    use std::sync::{Arc, Mutex};
432
433    #[test]
434    fn json_schema_request_preserves_output_schema() {
435        let schema = DirectJsonSchema {
436            name: "answer_shape".to_string(),
437            schema: json!({
438                "type": "object",
439                "properties": {
440                    "answer": { "type": "string" }
441                },
442                "required": ["answer"]
443            })
444            .into(),
445            strict: true,
446        };
447
448        let request = DirectRequest::json_schema("model-a", "return json", schema.clone());
449
450        assert_eq!(
451            request.output,
452            DirectOutputSpec::JsonSchema(schema),
453            "DirectRequest::json_schema must carry the requested output schema"
454        );
455    }
456
457    #[test]
458    fn direct_client_provider_accessors_expose_owned_provider_handle() {
459        let provider = TestProvider::builder()
460            .kind("direct-accessor-provider")
461            .serialize_config(|| json!({"provider": "owned"}))
462            .build()
463            .into_handle();
464        let mut client = DirectLlmClient::new(provider);
465
466        assert_eq!(client.provider().kind(), "direct-accessor-provider");
467        assert_eq!(
468            client.provider().to_spec().config,
469            json!({"provider": "owned"})
470        );
471
472        let options = ProviderOptions {
473            reliability: ProviderReliability::default().max_attempts(7),
474            max_output_tokens: Some(123),
475            ..Default::default()
476        };
477        client.provider_mut().set_options(options.clone());
478
479        assert_eq!(client.provider().options(), options);
480    }
481
482    #[tokio::test]
483    async fn direct_client_complete_delegates_to_provider_and_returns_response() {
484        let captured_request: Arc<Mutex<Option<LlmRequest>>> = Arc::new(Mutex::new(None));
485        let captured_for_provider = Arc::clone(&captured_request);
486        let provider = TestProvider::builder()
487            .kind("direct-complete-provider")
488            .complete(move |request| {
489                let captured_for_provider = Arc::clone(&captured_for_provider);
490                async move {
491                    *captured_for_provider.lock().expect("capture lock") = Some(request);
492                    Ok(LlmResponse {
493                        full_text: "provider delegated response".to_string(),
494                        parts: vec![LlmOutputPart::Text {
495                            text: "provider delegated response".to_string(),
496                            response_meta: None,
497                        }],
498                        usage: LlmUsage {
499                            input_tokens: 11,
500                            output_tokens: 3,
501                            ..Default::default()
502                        },
503                        terminal_reason: LlmTerminalReason::Stop,
504                        response_metadata: Default::default(),
505                        ..Default::default()
506                    })
507                }
508            })
509            .build()
510            .into_handle();
511        let mut client = DirectLlmClient::new(provider);
512        let mut request = DirectRequest::json("direct-model", "answer as json");
513        request.session_id = Some("direct-session".to_string());
514
515        let response = client
516            .complete(request)
517            .await
518            .expect("direct completion should delegate");
519
520        assert_eq!(response.full_text, "provider delegated response");
521        assert_eq!(response.llm_call.attempts.len(), 1);
522        let captured = captured_request
523            .lock()
524            .expect("capture lock")
525            .clone()
526            .expect("provider should receive a request");
527        assert_eq!(captured.model, "direct-model");
528        assert_eq!(captured.scope.session_id, "direct-session");
529        assert_eq!(captured.scope.agent_frame_id, "direct-session:frame:direct");
530        assert_eq!(captured.scope.request_id, "direct-session:direct");
531        assert!(matches!(
532            captured.output_spec,
533            Some(LlmOutputSpec::JsonObject)
534        ));
535        assert_eq!(captured.messages.len(), 1);
536    }
537
538    #[tokio::test]
539    async fn direct_client_validates_json_schema_output_against_canonical_schema() {
540        let provider = TestProvider::builder()
541            .kind("direct-validation-provider")
542            .complete(|_request| async {
543                Ok(LlmResponse {
544                    full_text: r#"{"items":[]}"#.to_string(),
545                    terminal_reason: LlmTerminalReason::Stop,
546                    response_metadata: Default::default(),
547                    ..Default::default()
548                })
549            })
550            .build()
551            .into_handle();
552        let mut client = DirectLlmClient::new(provider);
553        let request = DirectRequest::json_schema(
554            "direct-model",
555            "return items",
556            DirectJsonSchema {
557                name: "items_result".to_string(),
558                schema: json!({
559                    "type": "object",
560                    "required": ["items"],
561                    "properties": {
562                        "items": {
563                            "type": "array",
564                            "minItems": 1,
565                            "items": { "type": "string" }
566                        }
567                    }
568                })
569                .into(),
570                strict: true,
571            },
572        );
573
574        let err = client
575            .complete(request)
576            .await
577            .expect_err("empty items must fail canonical validation");
578
579        assert!(matches!(err, DirectLlmError::InvalidResponse(_)));
580        assert!(err.to_string().contains("items >= 1"));
581    }
582
583    fn reasoning_capability() -> ModelCapability {
584        ModelCapability {
585            reasoning: Some(crate::ReasoningCapability {
586                efforts: ["low", "medium", "high", "max"]
587                    .into_iter()
588                    .map(String::from)
589                    .collect(),
590                aliases: std::collections::BTreeMap::from([(
591                    "xhigh".to_string(),
592                    "max".to_string(),
593                )]),
594                ..Default::default()
595            }),
596            cache_control: None,
597            stream_termination: None,
598        }
599    }
600
601    #[tokio::test]
602    async fn direct_client_rejects_unsupported_effort_before_provider_call() {
603        let called = Arc::new(Mutex::new(false));
604        let called_for_provider = Arc::clone(&called);
605        let provider = TestProvider::builder()
606            .kind("direct-reject")
607            .complete(move |_request| {
608                let called = Arc::clone(&called_for_provider);
609                async move {
610                    *called.lock().expect("called lock") = true;
611                    Ok(LlmResponse::default())
612                }
613            })
614            .build()
615            .into_handle();
616        let mut client = DirectLlmClient::new(provider);
617
618        let mut request = DirectRequest::text("direct-model", "hi");
619        request.model_variant = crate::ReasoningSelection::Effort("turbo".to_string());
620        request.model_capability = reasoning_capability();
621
622        let err = client
623            .complete(request)
624            .await
625            .expect_err("unsupported effort must be rejected");
626        assert!(matches!(
627            err,
628            DirectLlmError::InvalidRequest {
629                category: ModelEffortValidationCategory::UnsupportedEffort,
630                ..
631            }
632        ));
633        assert!(err.to_string().contains("Unsupported effort `turbo`"));
634        assert!(
635            !*called.lock().expect("called lock"),
636            "the provider must not be called when the effort is rejected"
637        );
638    }
639
640    #[tokio::test]
641    async fn direct_client_normalizes_alias_effort_into_outgoing_request() {
642        let captured: Arc<Mutex<Option<crate::ReasoningSelection>>> = Arc::new(Mutex::new(None));
643        let captured_for_provider = Arc::clone(&captured);
644        let provider = TestProvider::builder()
645            .kind("direct-alias")
646            .complete(move |request| {
647                let captured = Arc::clone(&captured_for_provider);
648                async move {
649                    *captured.lock().expect("capture lock") = Some(request.model_variant.clone());
650                    Ok(LlmResponse {
651                        full_text: "ok".to_string(),
652                        terminal_reason: LlmTerminalReason::Stop,
653                        response_metadata: Default::default(),
654                        ..Default::default()
655                    })
656                }
657            })
658            .build()
659            .into_handle();
660        let mut client = DirectLlmClient::new(provider);
661
662        let mut request = DirectRequest::text("direct-model", "hi");
663        request.model_variant = crate::ReasoningSelection::Effort("XHigh".to_string());
664        request.model_capability = reasoning_capability();
665
666        client.complete(request).await.expect("completion");
667        let seen = captured
668            .lock()
669            .expect("capture lock")
670            .clone()
671            .expect("provider must be called");
672        assert_eq!(
673            seen,
674            crate::ReasoningSelection::Effort("max".to_string()),
675            "alias `XHigh` must clamp to canonical `max` before the provider sees the request"
676        );
677    }
678
679    #[tokio::test]
680    async fn direct_client_rejects_effort_when_model_is_not_configurable() {
681        let provider = TestProvider::builder()
682            .kind("direct-not-configurable")
683            .complete(|_request| async { Ok(LlmResponse::default()) })
684            .build()
685            .into_handle();
686        let mut client = DirectLlmClient::new(provider);
687
688        let mut request = DirectRequest::text("direct-model", "hi");
689        request.model_variant = crate::ReasoningSelection::Effort("high".to_string());
690        // No capability: the model exposes no configurable effort.
691
692        let err = client
693            .complete(request)
694            .await
695            .expect_err("effort on a non-configurable model must be rejected");
696        assert!(matches!(
697            err,
698            DirectLlmError::InvalidRequest {
699                category: ModelEffortValidationCategory::EffortNotConfigurable,
700                ..
701            }
702        ));
703    }
704
705    #[tokio::test]
706    async fn direct_client_rejects_missing_mandatory_effort() {
707        let provider = TestProvider::builder()
708            .kind("direct-mandatory")
709            .complete(|_request| async { Ok(LlmResponse::default()) })
710            .build()
711            .into_handle();
712        let mut client = DirectLlmClient::new(provider);
713
714        let mut capability = reasoning_capability();
715        capability.reasoning.as_mut().expect("reasoning").mandatory = true;
716        let mut request = DirectRequest::text("direct-model", "hi");
717        request.model_capability = capability;
718        // No model_variant supplied, but the model requires one.
719
720        let err = client
721            .complete(request)
722            .await
723            .expect_err("missing mandatory effort must be rejected");
724        assert!(matches!(
725            err,
726            DirectLlmError::InvalidRequest {
727                category: ModelEffortValidationCategory::EffortRequired,
728                ..
729            }
730        ));
731    }
732
733    #[test]
734    fn build_llm_request_preserves_nonempty_content_and_drops_empty_messages() {
735        let provider = TestProvider::default().into_handle();
736        let request = DirectRequest {
737            model: "input-model".to_string(),
738            messages: vec![
739                DirectMessage {
740                    role: DirectRole::System,
741                    parts: vec![DirectPart::Text(String::new())],
742                },
743                DirectMessage {
744                    role: DirectRole::User,
745                    parts: vec![
746                        DirectPart::Text("hello".to_string()),
747                        DirectPart::Text(String::new()),
748                    ],
749                },
750                DirectMessage {
751                    role: DirectRole::Assistant,
752                    parts: vec![DirectPart::Image(2)],
753                },
754            ],
755            attachments: Vec::new(),
756            output: DirectOutputSpec::Text,
757            generation: crate::GenerationOptions::default(),
758            stream_events: None,
759            session_id: None,
760            model_variant: Default::default(),
761            model_capability: ModelCapability::default(),
762            caused_by: None,
763            replay: None,
764        };
765
766        let llm_request = build_llm_request(&provider, request, "transport-model".to_string());
767
768        assert_eq!(llm_request.model, "transport-model");
769        assert_eq!(
770            llm_request.messages.len(),
771            2,
772            "empty normalized messages must be dropped"
773        );
774        assert_eq!(llm_request.messages[0].role, LlmRole::User);
775        assert_eq!(llm_request.messages[0].blocks.len(), 1);
776        assert!(matches!(
777            &llm_request.messages[0].blocks[0],
778            LlmContentBlock::Text { text, .. } if text.as_ref() == "hello"
779        ));
780        assert_eq!(llm_request.messages[1].role, LlmRole::Assistant);
781        assert!(matches!(
782            &llm_request.messages[1].blocks[0],
783            LlmContentBlock::Image { attachment_idx: 2 }
784        ));
785    }
786
787    #[test]
788    fn build_llm_request_preserves_direct_stream_sender_and_adds_required_noop_sender() {
789        let captured_events: Arc<Mutex<Vec<LlmStreamEvent>>> = Arc::new(Mutex::new(Vec::new()));
790        let captured_for_sender = Arc::clone(&captured_events);
791        let requested_sender = LlmEventSender::new(move |event| {
792            captured_for_sender
793                .lock()
794                .expect("stream event lock")
795                .push(event);
796        });
797        let mut request = DirectRequest::text("model", "prompt");
798        request.stream_events = Some(requested_sender);
799        let provider = TestProvider::default().into_handle();
800
801        let llm_request = build_llm_request(&provider, request, "model".to_string());
802        let sender = llm_request
803            .stream_events
804            .expect("explicit direct stream sender must be preserved");
805        sender.send(LlmStreamEvent::Delta("delta".to_string()));
806        assert_eq!(captured_events.lock().expect("stream event lock").len(), 1);
807
808        let streaming_provider = TestProvider::builder()
809            .requires_streaming(true)
810            .build()
811            .into_handle();
812        let llm_request = build_llm_request(
813            &streaming_provider,
814            DirectRequest::text("model", "prompt"),
815            "model".to_string(),
816        );
817        assert!(
818            llm_request.stream_events.is_some(),
819            "providers that require streaming need a no-op sender even when direct caller did not request one"
820        );
821    }
822}