use futures::Stream;
use reqwest::Method;
use serde_json::{Value, json};
use std::pin::Pin;
use tokio::time::timeout;
use crate::core::types::{
chat::ChatMessage,
chat::ChatRequest,
context::RequestContext,
message::MessageContent,
message::MessageRole,
responses::{ChatChoice, ChatChunk, ChatResponse, FinishReason},
};
use super::client::AzureAIClient;
use super::config::{AzureAIConfig, AzureAIEndpointType};
use crate::core::providers::base::{
HttpErrorMapper, SSETransformer, UnifiedSSEStream, read_streaming_error_body,
};
use crate::core::providers::unified_provider::ProviderError;
#[derive(Debug, Clone)]
pub struct AzureAIChatHandler {
client: AzureAIClient,
}
impl AzureAIChatHandler {
pub fn new(config: AzureAIConfig) -> Result<Self, ProviderError> {
Self::from_client(AzureAIClient::new(config)?)
}
pub(crate) fn from_client(client: AzureAIClient) -> Result<Self, ProviderError> {
Ok(Self { client })
}
pub async fn create_chat_completion(
&self,
request: ChatRequest,
_context: RequestContext,
) -> Result<ChatResponse, ProviderError> {
AzureAIChatUtils::validate_request(&request)?;
let azure_request = AzureAIChatUtils::transform_request(&request)?;
let url = self
.client
.get_config()
.build_endpoint_url(AzureAIEndpointType::ChatCompletions.as_path())
.map_err(|e| ProviderError::configuration("azure_ai", &e))?;
let response = self
.client
.request(Method::POST, &url)?
.json(&azure_request)
.send()
.await
.map_err(|e| ProviderError::network("azure_ai", format!("Request failed: {}", e)))?;
if !response.status().is_success() {
let status = response.status().as_u16();
let error_body = read_streaming_error_body(response)
.await
.map_err(|err| err.into_provider_error("azure_ai"))?;
return Err(HttpErrorMapper::map_status_code(
"azure_ai",
status,
&error_body,
));
}
let response_json: Value = response.json().await.map_err(|e| {
ProviderError::response_parsing("azure_ai", format!("Failed to parse response: {}", e))
})?;
AzureAIChatUtils::transform_response(response_json, &request.model)
}
pub async fn create_chat_completion_stream(
&self,
request: ChatRequest,
_context: RequestContext,
) -> Result<Pin<Box<dyn Stream<Item = Result<ChatChunk, ProviderError>> + Send>>, ProviderError>
{
AzureAIChatUtils::validate_request(&request)?;
let mut azure_request = AzureAIChatUtils::transform_request(&request)?;
azure_request["stream"] = json!(true);
let url = self
.client
.get_config()
.build_endpoint_url(AzureAIEndpointType::ChatCompletions.as_path())
.map_err(|e| ProviderError::configuration("azure_ai", &e))?;
let response = timeout(
self.client.get_config().timeout(),
self.client
.streaming_request(Method::POST, &url)?
.json(&azure_request)
.send(),
)
.await
.map_err(|_| ProviderError::timeout("azure_ai", "Streaming response header timeout"))?
.map_err(|error| ProviderError::network("azure_ai", error.to_string()))?;
if !response.status().is_success() {
let status = response.status().as_u16();
let error_body = read_streaming_error_body(response)
.await
.map_err(|err| err.into_provider_error("azure_ai"))?;
return Err(HttpErrorMapper::map_status_code(
"azure_ai",
status,
&error_body,
));
}
let transformer = AzureAISSETransformer::new(request.model.clone());
let stream = UnifiedSSEStream::new(Box::pin(response.bytes_stream()), transformer);
Ok(Box::pin(stream))
}
}
#[derive(Debug, Clone)]
struct AzureAISSETransformer {
model: String,
}
impl AzureAISSETransformer {
fn new(model: String) -> Self {
Self { model }
}
}
impl SSETransformer for AzureAISSETransformer {
fn provider_name(&self) -> &'static str {
"azure_ai"
}
fn transform_chunk(&self, data: &str) -> Result<Option<ChatChunk>, ProviderError> {
let chunk_data: Value = serde_json::from_str(data).map_err(|e| {
ProviderError::response_parsing("azure_ai", format!("Failed to parse SSE JSON: {}", e))
})?;
AzureAIChatUtils::transform_streaming_chunk(chunk_data, &self.model).map(Some)
}
}
pub struct AzureAIChatUtils;
impl AzureAIChatUtils {
pub fn validate_request(request: &ChatRequest) -> Result<(), ProviderError> {
if request.messages.is_empty() {
return Err(ProviderError::invalid_request(
"azure_ai",
"Messages cannot be empty",
));
}
if request.model.is_empty() {
return Err(ProviderError::invalid_request(
"azure_ai",
"Model cannot be empty",
));
}
if let Some(temp) = request.temperature
&& !(0.0..=2.0).contains(&temp)
{
return Err(ProviderError::invalid_request(
"azure_ai",
"Temperature must be between 0.0 and 2.0",
));
}
if let Some(top_p) = request.top_p
&& !(0.0..=1.0).contains(&top_p)
{
return Err(ProviderError::invalid_request(
"azure_ai",
"top_p must be between 0.0 and 1.0",
));
}
Ok(())
}
pub fn transform_request(request: &ChatRequest) -> Result<Value, ProviderError> {
let mut azure_request = json!({
"model": request.model,
"messages": Self::transform_messages(&request.messages)?
});
if let Some(temp) = request.temperature {
azure_request["temperature"] = json!(temp);
}
if let Some(max_tokens) = request.max_tokens {
azure_request["max_tokens"] = json!(max_tokens);
}
if let Some(max_completion_tokens) = request.max_completion_tokens {
azure_request["max_completion_tokens"] = json!(max_completion_tokens);
}
if let Some(top_p) = request.top_p {
azure_request["top_p"] = json!(top_p);
}
if let Some(freq_penalty) = request.frequency_penalty {
azure_request["frequency_penalty"] = json!(freq_penalty);
}
if let Some(pres_penalty) = request.presence_penalty {
azure_request["presence_penalty"] = json!(pres_penalty);
}
if let Some(stop) = &request.stop {
azure_request["stop"] = json!(stop);
}
if request.stream {
azure_request["stream"] = json!(true);
}
if let Some(tools) = &request.tools {
azure_request["tools"] = serde_json::to_value(tools).map_err(|e| {
ProviderError::transformation_error(
"azure_ai",
"request",
"azure_ai",
format!("Failed to serialize tools: {}", e),
)
})?;
}
if let Some(tool_choice) = &request.tool_choice {
azure_request["tool_choice"] = serde_json::to_value(tool_choice).map_err(|e| {
ProviderError::transformation_error(
"azure_ai",
"request",
"azure_ai",
format!("Failed to serialize tool_choice: {}", e),
)
})?;
}
Ok(azure_request)
}
fn transform_messages(messages: &[ChatMessage]) -> Result<Value, ProviderError> {
let mut azure_messages = Vec::new();
for message in messages {
let mut azure_message = json!({
"role": Self::transform_role(&message.role)
});
if let Some(content) = &message.content {
match content {
MessageContent::Text(text) => {
azure_message["content"] = json!(text);
}
MessageContent::Parts(parts) => {
let content_parts = parts
.iter()
.map(|part| {
json!(part)
})
.collect::<Vec<_>>();
azure_message["content"] = json!(content_parts);
}
}
}
if let Some(name) = &message.name {
azure_message["name"] = json!(name);
}
if let Some(function_call) = &message.function_call {
azure_message["function_call"] =
serde_json::to_value(function_call).map_err(|e| {
ProviderError::transformation_error(
"azure_ai",
"request",
"azure_ai",
format!("Failed to serialize function_call: {}", e),
)
})?;
}
if let Some(tool_calls) = &message.tool_calls {
azure_message["tool_calls"] = serde_json::to_value(tool_calls).map_err(|e| {
ProviderError::transformation_error(
"azure_ai",
"request",
"azure_ai",
format!("Failed to serialize tool_calls: {}", e),
)
})?;
}
if let Some(tool_call_id) = &message.tool_call_id {
azure_message["tool_call_id"] = json!(tool_call_id);
}
azure_messages.push(azure_message);
}
Ok(json!(azure_messages))
}
fn transform_role(role: &MessageRole) -> &'static str {
match role {
MessageRole::System => "system",
MessageRole::Developer => "developer",
MessageRole::User => "user",
MessageRole::Assistant => "assistant",
MessageRole::Function => "function",
MessageRole::Tool => "tool",
}
}
pub fn transform_response(response: Value, model: &str) -> Result<ChatResponse, ProviderError> {
let id = response["id"].as_str().unwrap_or("unknown").to_string();
let created = response["created"]
.as_i64()
.unwrap_or_else(|| chrono::Utc::now().timestamp());
let choices = response["choices"]
.as_array()
.ok_or_else(|| ProviderError::response_parsing("azure_ai", "Invalid choices format"))?
.iter()
.enumerate()
.map(|(index, choice)| Self::transform_choice(choice, index))
.collect::<Result<Vec<_>, _>>()?;
let usage = response
.get("usage")
.and_then(crate::core::providers::shared::strict_openai_chat_usage);
Ok(ChatResponse {
id,
object: "chat.completion".to_string(),
created,
model: model.to_string(),
choices,
usage,
system_fingerprint: response["system_fingerprint"]
.as_str()
.map(|s| s.to_string()),
})
}
fn transform_choice(choice: &Value, index: usize) -> Result<ChatChoice, ProviderError> {
let message_data = &choice["message"];
let role = match message_data["role"].as_str().unwrap_or("assistant") {
"system" => MessageRole::System,
"user" => MessageRole::User,
"assistant" => MessageRole::Assistant,
"function" => MessageRole::Function,
"tool" => MessageRole::Tool,
_ => MessageRole::Assistant,
};
let content = if let Some(content_str) = message_data["content"].as_str() {
MessageContent::Text(content_str.to_string())
} else {
MessageContent::Text(String::new())
};
let message = ChatMessage {
role,
content: Some(content),
thinking: None,
audio: None,
name: message_data["name"].as_str().map(|s| s.to_string()),
function_call: None, tool_calls: None, tool_call_id: message_data["tool_call_id"].as_str().map(|s| s.to_string()),
};
let finish_reason = match choice["finish_reason"].as_str() {
Some("stop") => Some(FinishReason::Stop),
Some("length") => Some(FinishReason::Length),
Some("content_filter") => Some(FinishReason::ContentFilter),
Some("tool_calls") => Some(FinishReason::ToolCalls),
Some("function_call") => Some(FinishReason::FunctionCall),
_ => None,
};
Ok(ChatChoice {
index: index as u32,
message,
finish_reason,
logprobs: None, })
}
pub fn parse_streaming_chunk(chunk_str: &str, model: &str) -> Result<ChatChunk, ProviderError> {
let lines: Vec<&str> = chunk_str.split("\n").collect();
for line in lines {
if let Some(data) = line.strip_prefix("data: ") {
if data == "[DONE]" {
return Ok(ChatChunk {
id: "stream_end".to_string(),
object: "chat.completion.chunk".to_string(),
created: chrono::Utc::now().timestamp(),
model: model.to_string(),
choices: vec![],
usage: None,
system_fingerprint: None,
});
}
let chunk_data: Value = serde_json::from_str(data).map_err(|e| {
ProviderError::response_parsing(
"azure_ai",
format!("Failed to parse chunk: {}", e),
)
})?;
return Self::transform_streaming_chunk(chunk_data, model);
}
}
Ok(ChatChunk {
id: "empty".to_string(),
object: "chat.completion.chunk".to_string(),
created: chrono::Utc::now().timestamp(),
model: model.to_string(),
choices: vec![],
usage: None,
system_fingerprint: None,
})
}
fn transform_streaming_chunk(
chunk_data: Value,
model: &str,
) -> Result<ChatChunk, ProviderError> {
let id = chunk_data["id"].as_str().unwrap_or("unknown").to_string();
let created = chunk_data["created"]
.as_i64()
.unwrap_or_else(|| chrono::Utc::now().timestamp());
let choices = if let Some(choices_array) = chunk_data["choices"].as_array() {
choices_array
.iter()
.enumerate()
.map(|(index, choice)| {
crate::core::types::responses::ChatStreamChoice {
index: index as u32,
delta: crate::core::types::responses::ChatDelta {
role: None,
content: choice["delta"]["content"].as_str().map(|s| s.to_string()),
thinking: None,
function_call: None,
tool_calls: None,
audio: None,
},
finish_reason: match choice["finish_reason"].as_str() {
Some("stop") => Some(FinishReason::Stop),
Some("length") => Some(FinishReason::Length),
Some("content_filter") => Some(FinishReason::ContentFilter),
Some("tool_calls") => Some(FinishReason::ToolCalls),
Some("function_call") => Some(FinishReason::FunctionCall),
_ => None,
},
logprobs: None,
}
})
.collect()
} else {
vec![]
};
Ok(ChatChunk {
id,
object: "chat.completion.chunk".to_string(),
created,
model: model.to_string(),
choices,
usage: None, system_fingerprint: None,
})
}
}
#[cfg(test)]
#[path = "chat_tests.rs"]
mod tests;