1use bytes::Bytes;
4use futures::{Stream, StreamExt};
5use reqwest::Client;
6use serde::Deserialize;
7use serde_json::Value as JsonValue;
8use std::future::Future;
9use std::pin::Pin;
10use std::sync::Arc;
11
12use crate::{
13 Api, AssistantMessage, ContentBlock, Context, Model, Provider, ProviderEvent, StopReason,
14 StreamOptions, StreamResult, Usage, error::ProviderError,
15};
16
17use super::shared_client;
18
19#[derive(Clone)]
34pub struct AzureProvider {
35 client: &'static Client,
36 api_key: Option<String>,
37 resource_name: Option<String>,
38 deployment_name: Option<String>,
39}
40
41impl AzureProvider {
42 pub fn new() -> Self {
46 Self {
47 client: shared_client(),
48 api_key: None,
49 resource_name: None,
50 deployment_name: None,
51 }
52 }
53
54 #[cfg(test)]
56 pub fn with_config(
57 api_key: impl Into<String>,
58 resource_name: impl Into<String>,
59 deployment_name: impl Into<String>,
60 ) -> Self {
61 Self {
62 client: shared_client(),
63 api_key: Some(api_key.into()),
64 resource_name: Some(resource_name.into()),
65 deployment_name: Some(deployment_name.into()),
66 }
67 }
68
69 fn build_url(&self, model: &Model) -> Result<String, ProviderError> {
71 if !model.base_url.is_empty() && model.base_url != "https://api.openai.com" {
73 return Ok(format!(
75 "{}/chat/completions?api-version=2024-02-15-preview",
76 model.base_url.trim_end_matches('/')
77 ));
78 }
79
80 let resource = self.resource_name.as_ref().ok_or_else(|| {
82 ProviderError::InvalidResponse("AZURE_OPENAI_RESOURCE_NAME not set".into())
83 })?;
84
85 let deployment = self.deployment_name.as_ref().ok_or_else(|| {
86 ProviderError::InvalidResponse("AZURE_OPENAI_DEPLOYMENT_NAME not set".into())
87 })?;
88
89 let url = format!(
90 "https://{}.openai.azure.com/openai/deployments/{}/chat/completions?api-version=2024-02-15-preview",
91 resource, deployment
92 );
93
94 Ok(url)
95 }
96
97 fn get_api_key(&self, options: &Option<StreamOptions>) -> Result<String, ProviderError> {
99 options
100 .as_ref()
101 .and_then(|o| o.api_key.as_ref())
102 .or(self.api_key.as_ref())
103 .cloned()
104 .ok_or_else(|| ProviderError::MissingApiKey)
105 }
106
107 fn build_headers(
109 &self,
110 api_key: &str,
111 options: &Option<StreamOptions>,
112 ) -> Result<reqwest::header::HeaderMap, ProviderError> {
113 let mut headers = reqwest::header::HeaderMap::new();
114
115 headers.insert(
117 "api-key",
118 api_key.parse().map_err(|e| {
119 ProviderError::InvalidResponse(format!("invalid header value: {e}"))
120 })?,
121 );
122 headers.insert(
123 reqwest::header::CONTENT_TYPE,
124 "application/json".parse().map_err(|e| {
125 ProviderError::InvalidResponse(format!("invalid header value: {e}"))
126 })?,
127 );
128
129 if let Some(opts) = options {
131 for (k, v) in &opts.headers {
132 if let (Ok(name), Ok(value)) = (
133 k.parse::<reqwest::header::HeaderName>(),
134 v.parse::<reqwest::header::HeaderValue>(),
135 ) {
136 headers.insert(name, value);
137 }
138 }
139 }
140
141 Ok(headers)
142 }
143}
144
145impl Default for AzureProvider {
146 fn default() -> Self {
147 Self::new()
148 }
149}
150
151impl Provider for AzureProvider {
152 fn stream<'a>(
153 &'a self,
154 model: &'a Model,
155 context: &'a Context,
156 options: Option<StreamOptions>,
157 ) -> Pin<Box<dyn Future<Output = StreamResult> + Send + 'a>> {
158 Box::pin(async move {
159 let url = self.build_url(model)?;
161
162 let api_key = self.get_api_key(&options)?;
164
165 let messages = build_messages(context)?;
167
168 let mut body = serde_json::json!({
170 "messages": messages,
171 "stream": true,
172 });
173
174 if model.id != "default" && model.id != "azure" {
176 body["model"] = serde_json::json!(model.id);
177 }
178
179 if let Some(ref opts) = options {
181 if let Some(temp) = opts.temperature {
182 body["temperature"] = serde_json::json!(temp);
183 }
184
185 if let Some(max) = opts.max_tokens {
186 body["max_tokens"] = serde_json::json!(max);
187 }
188 }
189
190 if !context.tools.is_empty() {
192 body["tools"] = build_tools(&context.tools)?;
193 }
194
195 let headers = self.build_headers(&api_key, &options)?;
197
198 let response = self
200 .client
201 .post(&url)
202 .headers(headers)
203 .json(&body)
204 .send()
205 .await
206 .map_err(ProviderError::RequestFailed)?;
207
208 if !response.status().is_success() {
209 let status = response.status();
210 let body: String = response.text().await.unwrap_or_default();
211 return Err(ProviderError::HttpError(
212 crate::error::HttpErrorDetail::new(status.as_u16(), body),
213 ));
214 }
215
216 let provider_name = model.provider.clone();
218 let model_id = model.id.clone();
219
220 let stream = response
221 .bytes_stream()
222 .scan(
223 Vec::<u8>::new(),
224 move |pending_bytes, chunk: Result<Bytes, reqwest::Error>| {
225 let pn = provider_name.clone();
226 let mid = model_id.clone();
227 let events = match chunk {
232 Ok(bytes) => {
233 let mut combined =
238 Vec::with_capacity(pending_bytes.len() + bytes.len());
239 combined.extend_from_slice(pending_bytes);
240 combined.extend_from_slice(&bytes);
241 let (text, trailing) = super::sse::split_complete_lines(&combined);
242 *pending_bytes = trailing;
243 parse_sse_events(&text, &pn, &mid)
244 }
245 Err(e) => vec![ProviderEvent::Error {
246 reason: StopReason::Error,
247 error: create_error_message(&e.to_string(), &pn, &mid),
248 }],
249 };
250 std::future::ready(Some(futures::stream::iter(events)))
251 },
252 )
253 .flatten();
254
255 Ok(Box::pin(stream) as Pin<Box<dyn Stream<Item = ProviderEvent> + Send>>)
256 })
257 }
258}
259
260fn build_messages(context: &Context) -> Result<Vec<JsonValue>, ProviderError> {
262 let mut messages = Vec::new();
263
264 if let Some(ref prompt) = context.system_prompt {
266 messages.push(serde_json::json!({
267 "role": "system",
268 "content": prompt,
269 }));
270 }
271
272 for msg in &context.messages {
274 match msg {
275 crate::Message::User(u) => {
276 let content: String = match &u.content {
277 crate::MessageContent::Text(s) => s.clone(),
278 crate::MessageContent::Blocks(blocks) => blocks_to_content(blocks)?.to_string(),
279 };
280 messages.push(serde_json::json!({
281 "role": "user",
282 "content": content,
283 }));
284 }
285 crate::Message::Assistant(a) => {
286 let content = blocks_to_content(&a.content)?.to_string();
287 messages.push(serde_json::json!({
288 "role": "assistant",
289 "content": content,
290 }));
291 }
292 crate::Message::ToolResult(t) => {
293 let content = blocks_to_content(&t.content)?.to_string();
294 messages.push(serde_json::json!({
295 "role": "tool",
296 "tool_call_id": t.tool_call_id,
297 "tool_name": t.tool_name,
298 "content": content,
299 }));
300 }
301 }
302 }
303
304 Ok(messages)
305}
306
307fn blocks_to_content(blocks: &[ContentBlock]) -> Result<JsonValue, ProviderError> {
309 if blocks.len() == 1
310 && let Some(text) = blocks[0].as_text()
311 {
312 return Ok(JsonValue::String(text.to_string()));
313 }
314
315 let items: Result<Vec<_>, _> = blocks
316 .iter()
317 .map(|block| match block {
318 ContentBlock::Text(t) => Ok(serde_json::json!({
319 "type": "text",
320 "text": t.text,
321 })),
322 ContentBlock::ToolCall(tc) => Ok(serde_json::json!({
323 "type": "function",
324 "id": tc.id,
325 "function": {
326 "name": tc.name,
327 "arguments": tc.arguments.to_string(),
328 },
329 })),
330 ContentBlock::Thinking(th) => Ok(serde_json::json!({
331 "type": "thinking",
332 "thinking": th.thinking,
333 })),
334 ContentBlock::Image(img) => Ok(serde_json::json!({
335 "type": "image_url",
336 "image_url": {
337 "url": format!("data:{};base64,{}", img.mime_type, img.data),
338 },
339 })),
340 ContentBlock::Unknown(_) => Err(ProviderError::InvalidResponse(
341 "Unknown content block type".into(),
342 )),
343 })
344 .collect();
345
346 Ok(serde_json::json!(items?))
347}
348
349fn build_tools(tools: &[crate::Tool]) -> Result<JsonValue, ProviderError> {
351 let items: Vec<_> = tools
352 .iter()
353 .map(|tool| {
354 serde_json::json!({
355 "type": "function",
356 "function": {
357 "name": tool.name,
358 "description": tool.description,
359 "parameters": tool.parameters,
360 },
361 })
362 })
363 .collect();
364
365 Ok(serde_json::json!(items))
366}
367
368fn parse_sse_events(text: &str, provider: &str, model_id: &str) -> Vec<ProviderEvent> {
372 let mut events = Vec::with_capacity(text.len() / 80);
373 let mut partial_message = AssistantMessage::new(Api::OpenAiCompletions, provider, model_id);
374
375 let mut accumulated_usage = Usage::default();
376
377 for line in text.split('\n') {
378 let line = line.trim_end_matches('\r');
379 if line.is_empty() {
380 continue;
381 }
382
383 if !line.starts_with("data: ") {
385 continue;
386 }
387
388 let data = &line[6..]; if data == "[DONE]" {
392 break;
393 }
394
395 if data.is_empty() {
396 continue;
397 }
398
399 let chunk = match serde_json::from_str::<SSEChunk>(data) {
400 Ok(c) => c,
401 Err(_) => continue,
402 };
403
404 let this_chunk_usage = chunk.usage.as_ref();
406
407 for choice in &chunk.choices {
408 if let Some(delta) = &choice.delta {
409 if let Some(content) = &delta.content {
410 let last_text_idx = partial_message
413 .content
414 .iter()
415 .rposition(|b| matches!(b, ContentBlock::Text(_)));
416 if let Some(idx) = last_text_idx
417 && let ContentBlock::Text(t) = &mut partial_message.content[idx]
418 {
419 t.text.push_str(content);
420 } else {
421 partial_message
422 .content
423 .push(ContentBlock::Text(crate::TextContent::new(content.clone())));
424 }
425 events.push(ProviderEvent::TextDelta {
426 content_index: choice.index,
427 delta: content.clone(),
428 partial: Arc::new(partial_message.clone()),
429 });
430 }
431
432 if let Some(tool_calls) = &delta.tool_calls {
433 for tc in tool_calls {
434 if let Some(func) = &tc.function {
435 events.push(ProviderEvent::ToolCallDelta {
436 content_index: choice.index,
437 delta: func.arguments.clone().unwrap_or_default(),
438 partial: Arc::new(partial_message.clone()),
439 });
440 }
441 }
442 }
443 }
444
445 if choice.finish_reason.is_some() {
446 let reason = match choice.finish_reason.as_deref() {
449 Some("stop") => StopReason::Stop,
450 Some("length") => StopReason::Length,
451 Some("tool_calls") => StopReason::ToolUse,
452 _ => StopReason::Stop,
453 };
454
455 let mut done_msg = partial_message.clone();
456
457 if let Some(usage) = this_chunk_usage {
459 done_msg.usage.input = usage.prompt_tokens;
460 done_msg.usage.output = usage.completion_tokens;
461 done_msg.usage.cache_read = usage
462 .prompt_tokens_details
463 .as_ref()
464 .map(|d| d.cached_tokens)
465 .unwrap_or(0);
466 done_msg.usage.total_tokens = usage.total_tokens;
467 } else {
468 done_msg.usage = accumulated_usage.clone();
469 }
470
471 events.push(ProviderEvent::Done {
472 reason,
473 message: done_msg,
474 });
475 }
476 }
477
478 if let Some(usage) = this_chunk_usage {
480 accumulated_usage.input = usage.prompt_tokens;
481 accumulated_usage.output = usage.completion_tokens;
482 accumulated_usage.cache_read = usage
483 .prompt_tokens_details
484 .as_ref()
485 .map(|d| d.cached_tokens)
486 .unwrap_or(0);
487 accumulated_usage.total_tokens = usage.total_tokens;
488 }
489 }
490
491 events
492}
493
494fn create_error_message(msg: &str, provider: &str, model_id: &str) -> AssistantMessage {
496 let mut message = AssistantMessage::new(Api::OpenAiCompletions, provider, model_id);
497 message.stop_reason = StopReason::Error;
498 message.error_message = Some(msg.to_string());
499 message
500}
501
502#[derive(Debug, Deserialize)]
504struct SSEChunk {
505 _id: Option<String>,
506 #[serde(rename = "model")]
507 _model: Option<String>,
508 choices: Vec<Choice>,
509 usage: Option<UsageInfo>,
510}
511
512#[derive(Debug, Deserialize)]
513struct Choice {
514 index: usize,
515 delta: Option<Delta>,
516 finish_reason: Option<String>,
517}
518
519#[derive(Debug, Deserialize)]
520struct Delta {
521 content: Option<String>,
522 tool_calls: Option<Vec<ToolCallDelta>>,
523}
524
525#[derive(Debug, Deserialize)]
526struct ToolCallDelta {
527 _index: Option<usize>,
528 _id: Option<String>,
529 #[serde(rename = "type")]
530 _type_: Option<String>,
531 function: Option<FunctionDelta>,
532}
533
534#[derive(Debug, Deserialize)]
535struct FunctionDelta {
536 _name: Option<String>,
537 arguments: Option<String>,
538}
539
540#[derive(Debug, Deserialize, Clone)]
541struct UsageInfo {
542 prompt_tokens: usize,
543 completion_tokens: usize,
544 total_tokens: usize,
545 #[serde(rename = "prompt_tokens_details")]
546 prompt_tokens_details: Option<PromptTokensDetails>,
547}
548
549#[derive(Debug, Deserialize, Clone)]
550struct PromptTokensDetails {
551 #[serde(rename = "cached_tokens")]
552 cached_tokens: usize,
553}
554
555#[cfg(test)]
556mod tests {
557 use super::*;
558
559 fn make_test_model(id: &str, base_url: &str) -> Model {
560 Model::new(id, id, Api::OpenAiCompletions, "azure", base_url)
561 }
562
563 #[test]
564 fn test_build_url_from_base_url() {
565 let provider = AzureProvider::new();
566 let model = make_test_model(
567 "gpt-4o",
568 "https://my-resource.openai.azure.com/openai/deployments/gpt-4o",
569 );
570
571 let url = provider.build_url(&model).unwrap();
572 assert!(url.contains("api-version=2024-02-15-preview"));
573 assert!(url.contains("my-resource"));
574 assert!(url.contains("gpt-4o"));
575 }
576
577 #[test]
578 fn test_build_url_missing_resource() {
579 let provider = AzureProvider {
580 client: shared_client(),
581 api_key: Some("test-key".to_string()),
582 resource_name: None,
583 deployment_name: Some("gpt-4o".to_string()),
584 };
585
586 let model = make_test_model("default", "");
587
588 let result = provider.build_url(&model);
589 assert!(result.is_err());
590 match result.unwrap_err() {
591 ProviderError::InvalidResponse(msg) => {
592 assert!(msg.contains("AZURE_OPENAI_RESOURCE_NAME"));
593 }
594 _ => panic!("Expected InvalidResponse"),
595 }
596 }
597
598 #[test]
599 fn test_build_url_missing_deployment() {
600 let provider = AzureProvider {
601 client: shared_client(),
602 api_key: Some("test-key".to_string()),
603 resource_name: Some("my-resource".to_string()),
604 deployment_name: None,
605 };
606
607 let model = make_test_model("default", "");
608
609 let result = provider.build_url(&model);
610 assert!(result.is_err());
611 match result.unwrap_err() {
612 ProviderError::InvalidResponse(msg) => {
613 assert!(msg.contains("AZURE_OPENAI_DEPLOYMENT_NAME"));
614 }
615 _ => panic!("Expected InvalidResponse"),
616 }
617 }
618
619 #[test]
620 fn test_build_url_from_env_vars() {
621 let provider = AzureProvider {
622 client: shared_client(),
623 api_key: Some("test-key".to_string()),
624 resource_name: Some("my-resource".to_string()),
625 deployment_name: Some("gpt-4o".to_string()),
626 };
627
628 let model = make_test_model("default", "");
629
630 let url = provider.build_url(&model).unwrap();
631 assert_eq!(
632 url,
633 "https://my-resource.openai.azure.com/openai/deployments/gpt-4o/chat/completions?api-version=2024-02-15-preview"
634 );
635 }
636
637 #[test]
638 fn test_parse_sse_events_text() {
639 let sse_data = r#"data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4o","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}
640
641data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4o","choices":[{"index":0,"delta":{"content":" world"},"finish_reason":"stop"}]}
642
643data: [DONE]"#;
644
645 let events = parse_sse_events(sse_data, "azure", "gpt-4o");
646
647 assert!(events.len() >= 3);
649
650 match &events[0] {
652 ProviderEvent::TextDelta { delta, .. } => assert_eq!(delta, "Hello"),
653 _ => panic!("Expected TextDelta event"),
654 }
655
656 match &events[events.len() - 1] {
658 ProviderEvent::Done { reason, .. } => assert_eq!(*reason, StopReason::Stop),
659 _ => panic!("Expected Done event"),
660 }
661 }
662
663 #[test]
664 fn test_parse_sse_events_with_tool_calls() {
665 let sse_data = r#"data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4o","choices":[{"index":0,"delta":{"tool_calls":[{"id":"call_abc123","type":"function","function":{"name":"get_weather","arguments":""}}]},"finish_reason":null}]}
666
667data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4o","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"location\":"}}]},"finish_reason":null}]}
668
669data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4o","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"Boston\"}"}}]},"finish_reason":"tool_calls"}]}
670
671data: [DONE]"#;
672
673 let events = parse_sse_events(sse_data, "azure", "gpt-4o");
674
675 assert!(events.len() >= 4);
677
678 let has_tool_call = events
680 .iter()
681 .any(|e| matches!(e, ProviderEvent::ToolCallDelta { .. }));
682 assert!(
683 has_tool_call,
684 "Should have at least one ToolCallDelta event"
685 );
686
687 match &events[events.len() - 1] {
689 ProviderEvent::Done { reason, .. } => assert_eq!(*reason, StopReason::ToolUse),
690 _ => panic!("Expected Done event with ToolUse reason"),
691 }
692 }
693
694 #[test]
695 fn test_parse_sse_events_usage() {
696 let sse_data = r#"data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4o","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":"stop"}],"usage":{"prompt_tokens":10,"completion_tokens":5,"total_tokens":15,"prompt_tokens_details":{"cached_tokens":0}}}
697
698data: [DONE]"#;
699
700 let events = parse_sse_events(sse_data, "azure", "gpt-4o");
701
702 let done_event = events
704 .iter()
705 .find(|e| matches!(e, ProviderEvent::Done { .. }));
706 assert!(done_event.is_some());
707
708 if let ProviderEvent::Done { message, .. } = done_event.unwrap() {
709 assert_eq!(message.usage.input, 10);
710 assert_eq!(message.usage.output, 5);
711 assert_eq!(message.usage.total_tokens, 15);
712 }
713 }
714
715 #[test]
716 fn test_build_headers_includes_api_key() {
717 let provider = AzureProvider::new();
718 let api_key = "test-api-key-12345";
719
720 let headers = provider.build_headers(api_key, &None).unwrap();
721
722 let api_key_header = headers.get("api-key");
724 assert!(api_key_header.is_some());
725 assert_eq!(api_key_header.unwrap().to_str().unwrap(), api_key);
726
727 let content_type = headers.get(reqwest::header::CONTENT_TYPE);
729 assert!(content_type.is_some());
730 }
731
732 #[test]
733 fn test_build_headers_no_bearer_token() {
734 let provider = AzureProvider::new();
735 let api_key = "test-api-key-12345";
736
737 let headers = provider.build_headers(api_key, &None).unwrap();
738
739 let auth_header = headers.get(reqwest::header::AUTHORIZATION);
741 assert!(
742 auth_header.is_none(),
743 "Azure should not use Bearer token authentication"
744 );
745 }
746
747 #[test]
748 fn test_with_config_constructor() {
749 let provider = AzureProvider::with_config("my-api-key", "my-resource", "gpt-4o");
750
751 let model = make_test_model("default", "");
753
754 let url = provider.build_url(&model).unwrap();
755 assert!(url.contains("my-resource"));
756 assert!(url.contains("gpt-4o"));
757 }
758
759 #[test]
760 fn test_azure_endpoint_format() {
761 let provider = AzureProvider {
762 client: shared_client(),
763 api_key: Some("key".to_string()),
764 resource_name: Some("my-resource".to_string()),
765 deployment_name: Some("gpt-4-turbo".to_string()),
766 };
767
768 let model = make_test_model("default", "");
769 let url = provider.build_url(&model).unwrap();
770
771 assert!(url.starts_with("https://"));
773 assert!(url.contains(".openai.azure.com"));
774 assert!(url.contains("/openai/deployments/"));
775 assert!(url.contains("gpt-4-turbo"));
776 assert!(url.contains("chat/completions"));
777 assert!(url.contains("api-version=2024-02-15-preview"));
778 }
779}