refactor: move bedrock proxy handling to windmill-ai (#9309)

* refactor: move bedrock proxy handling to windmill-ai

* docs: track ai refactor follow-ups
This commit is contained in:
centdix
2026-06-08 11:35:45 +02:00
committed by tristantr
parent 528a7afcd0
commit 2df3dd09dd
9 changed files with 923 additions and 998 deletions
+1 -5
View File
@@ -13875,6 +13875,7 @@ dependencies = [
"async-trait",
"aws-config",
"aws-credential-types",
"aws-sdk-bedrock",
"aws-sdk-bedrockruntime",
"aws-smithy-types",
"base64 0.22.1",
@@ -13923,13 +13924,8 @@ dependencies = [
"async-stream",
"async-trait",
"async_zip",
"aws-config",
"aws-credential-types",
"aws-sdk-bedrock",
"aws-sdk-bedrockruntime",
"aws-sdk-config",
"aws-sigv4",
"aws-smithy-types",
"axum 0.8.9",
"base32",
"base64 0.22.1",
+2 -1
View File
@@ -6,7 +6,7 @@ edition.workspace = true
[features]
default = []
bedrock = ["dep:aws-sdk-bedrockruntime", "dep:aws-credential-types", "dep:aws-smithy-types", "dep:aws-config"]
bedrock = ["dep:aws-sdk-bedrock", "dep:aws-sdk-bedrockruntime", "dep:aws-credential-types", "dep:aws-smithy-types", "dep:aws-config"]
mcp = ["dep:windmill-mcp"]
[lib]
@@ -42,4 +42,5 @@ ulid.workspace = true
aws-config = { workspace = true, optional = true }
aws-credential-types = { workspace = true, optional = true }
aws-smithy-types = { workspace = true, optional = true }
aws-sdk-bedrock = { workspace = true, optional = true }
aws-sdk-bedrockruntime = { workspace = true, optional = true }
+775 -8
View File
@@ -7,21 +7,731 @@
//! - Helper utilities
use crate::{
ai_bedrock::{
bedrock_model_supports_prompt_caching, bedrock_stream_event_is_block_stop,
bedrock_stream_event_to_text, bedrock_stream_event_to_tool_delta,
bedrock_stream_event_to_tool_start, build_tool_config, create_inference_config,
format_bedrock_error, openai_messages_to_bedrock, streaming_tool_calls_to_openai,
BearerTokenProvider, BedrockClient, StreamingToolCall,
},
ai_providers::USE_ENV_REGION,
ai_types::{OpenAIFunction, OpenAIToolCall, ToolDefFunction},
image_handler::prepare_messages_for_api,
proxy::ProxyBuildArgs,
query_builder::{ParsedResponse, StreamEventSink},
types::{OpenAIMessage, StreamingEvent, TokenUsage, ToolDef},
};
use bytes::Bytes;
use futures::{stream::BoxStream, StreamExt};
use http::{HeaderMap, Method, StatusCode};
use serde::Deserialize;
use std::collections::HashMap;
use windmill_common::{client::AuthedClient, error::Error};
// Import shared Bedrock helpers for provider orchestration.
use crate::ai_bedrock::{
bedrock_model_supports_prompt_caching, bedrock_stream_event_is_block_stop,
bedrock_stream_event_to_text, bedrock_stream_event_to_tool_delta,
bedrock_stream_event_to_tool_start, build_tool_config, create_inference_config,
format_bedrock_error, openai_messages_to_bedrock, streaming_tool_calls_to_openai,
BedrockClient, StreamingToolCall,
};
// ============================================================================
// Native Proxy Execution
// ============================================================================
/// OpenAI-format request body for Bedrock SDK proxy handlers.
#[derive(Deserialize, Debug)]
struct OpenAIRequest {
messages: Vec<OpenAIMessage>,
#[serde(default)]
tools: Option<Vec<OpenAIToolDef>>,
#[serde(default)]
tool_choice: Option<serde_json::Value>,
#[serde(default)]
max_tokens: Option<i32>,
#[serde(default)]
temperature: Option<f32>,
}
#[derive(Deserialize, Debug)]
struct OpenAIToolDef {
#[serde(default)]
#[allow(dead_code)]
r#type: Option<String>,
function: OpenAIToolFunction,
}
#[derive(Deserialize, Debug)]
struct OpenAIToolFunction {
name: String,
#[serde(default)]
description: Option<String>,
#[serde(default)]
parameters: Option<serde_json::Value>,
}
#[derive(Deserialize, Debug)]
struct BedrockProxyChatRequest {
model: String,
#[serde(default)]
stream: bool,
}
enum BedrockAuthConfig {
BearerToken(String),
IamCredentials {
access_key_id: String,
secret_access_key: String,
session_token: Option<String>,
},
Environment,
}
pub enum BedrockProxyResponseBody {
Fixed(Bytes),
Stream(BoxStream<'static, std::result::Result<Bytes, std::io::Error>>),
}
pub struct BedrockProxyResponse {
pub status_code: StatusCode,
pub headers: HeaderMap,
pub body: BedrockProxyResponseBody,
}
/// Handle a workspace Bedrock proxy request through the AWS SDK.
///
/// The API still owns credential resolution, route authorization, auditing, and
/// cache behavior. This helper owns Bedrock-specific control-plane and
/// OpenAI-compatible Converse transformations.
pub async fn handle_bedrock_proxy(
args: &ProxyBuildArgs<'_>,
) -> Result<BedrockProxyResponse, Error> {
let region = args.credentials.region.as_deref().unwrap_or(USE_ENV_REGION);
if *args.method == Method::GET {
return match args.path {
"foundation-models" => list_foundation_models(args, region).await,
"inference-profiles" => list_inference_profiles(args, region).await,
_ => Err(Error::BadRequest(format!(
"Unsupported AWS Bedrock proxy path: {}",
args.path
))),
};
}
if *args.method != Method::POST {
return Err(Error::BadRequest(format!(
"Unsupported AWS Bedrock proxy method: {}",
args.method
)));
}
let request: BedrockProxyChatRequest = serde_json::from_slice(args.body)
.map_err(|e| Error::internal_err(format!("Failed to parse request body: {}", e)))?;
if request.stream {
handle_bedrock_sdk_streaming(&request.model, args.body, args, region).await
} else {
handle_bedrock_sdk_non_streaming(&request.model, args.body, args, region).await
}
}
fn determine_auth_config(
api_key: Option<&str>,
aws_access_key_id: Option<&str>,
aws_secret_access_key: Option<&str>,
aws_session_token: Option<&str>,
) -> BedrockAuthConfig {
if let Some(key) = api_key.filter(|k| !k.is_empty()) {
BedrockAuthConfig::BearerToken(key.to_string())
} else if let (Some(access_key_id), Some(secret_access_key)) = (
aws_access_key_id.filter(|s| !s.is_empty()),
aws_secret_access_key.filter(|s| !s.is_empty()),
) {
BedrockAuthConfig::IamCredentials {
access_key_id: access_key_id.to_string(),
secret_access_key: secret_access_key.to_string(),
session_token: aws_session_token
.filter(|token| !token.is_empty())
.map(str::to_string),
}
} else {
BedrockAuthConfig::Environment
}
}
async fn create_bedrock_client(
args: &ProxyBuildArgs<'_>,
region: &str,
) -> Result<BedrockClient, Error> {
match determine_auth_config(
args.credentials.api_key.as_deref(),
args.credentials.aws_access_key_id.as_deref(),
args.credentials.aws_secret_access_key.as_deref(),
args.credentials.aws_session_token.as_deref(),
) {
BedrockAuthConfig::BearerToken(key) => BedrockClient::from_bearer_token(key, region).await,
BedrockAuthConfig::IamCredentials { access_key_id, secret_access_key, session_token } => {
BedrockClient::from_credentials(access_key_id, secret_access_key, session_token, region)
.await
}
BedrockAuthConfig::Environment => BedrockClient::from_env(region).await,
}
}
fn build_tool_config_from_request(
tools: Option<&[OpenAIToolDef]>,
tool_choice: Option<&serde_json::Value>,
enable_prompt_caching: bool,
) -> Result<Option<aws_sdk_bedrockruntime::types::ToolConfiguration>, Error> {
if let Some(tools) = tools {
let tool_defs: Vec<ToolDef> = tools
.iter()
.map(|t| ToolDef {
r#type: "function".to_string(),
function: ToolDefFunction {
name: t.function.name.clone(),
description: t.function.description.clone(),
parameters: Box::from(
serde_json::value::RawValue::from_string(
serde_json::to_string(
&t.function
.parameters
.clone()
.unwrap_or(serde_json::json!({})),
)
.unwrap_or_default(),
)
.unwrap_or_else(|_| {
serde_json::value::RawValue::from_string("{}".to_string()).unwrap()
}),
),
},
})
.collect();
let force_tool_use = tool_choice
.map(|tc| tc == "required" || tc.as_str() == Some("required"))
.unwrap_or(false);
build_tool_config(Some(&tool_defs), force_tool_use, enable_prompt_caching)
} else {
Ok(None)
}
}
async fn create_bedrock_control_client(
args: &ProxyBuildArgs<'_>,
region: &str,
) -> Result<aws_sdk_bedrock::Client, Error> {
use aws_config::BehaviorVersion;
let region_provider = aws_sdk_bedrock::config::Region::new(region.to_string());
match determine_auth_config(
args.credentials.api_key.as_deref(),
args.credentials.aws_access_key_id.as_deref(),
args.credentials.aws_secret_access_key.as_deref(),
args.credentials.aws_session_token.as_deref(),
) {
BedrockAuthConfig::BearerToken(key) => {
let config = aws_sdk_bedrock::config::Builder::new()
.region(region_provider)
.behavior_version(BehaviorVersion::latest())
.token_provider(BearerTokenProvider::new(key))
.build();
Ok(aws_sdk_bedrock::Client::from_conf(config))
}
BedrockAuthConfig::IamCredentials { access_key_id, secret_access_key, session_token } => {
let credentials = aws_credential_types::Credentials::new(
access_key_id,
secret_access_key,
session_token,
None,
"windmill",
);
let config = aws_sdk_bedrock::config::Builder::new()
.region(region_provider)
.behavior_version(BehaviorVersion::latest())
.credentials_provider(credentials)
.build();
Ok(aws_sdk_bedrock::Client::from_conf(config))
}
BedrockAuthConfig::Environment => {
let config = aws_config::defaults(BehaviorVersion::latest())
.region(region_provider)
.load()
.await;
Ok(aws_sdk_bedrock::Client::new(&config))
}
}
}
async fn list_foundation_models(
args: &ProxyBuildArgs<'_>,
region: &str,
) -> Result<BedrockProxyResponse, Error> {
let client = create_bedrock_control_client(args, region).await?;
let response = client
.list_foundation_models()
.send()
.await
.map_err(|e| Error::internal_err(format!("Failed to list foundation models: {}", e)))?;
let models: Vec<serde_json::Value> = response
.model_summaries()
.iter()
.map(|m| {
serde_json::json!({
"modelId": m.model_id(),
"modelName": m.model_name(),
"providerName": m.provider_name(),
"modelArn": m.model_arn(),
"inputModalities": m.input_modalities().iter().map(|i| i.as_str()).collect::<Vec<_>>(),
"outputModalities": m.output_modalities().iter().map(|o| o.as_str()).collect::<Vec<_>>(),
"responseStreamingSupported": m.response_streaming_supported(),
"inferenceTypesSupported": m.inference_types_supported().iter().map(|i| i.as_str()).collect::<Vec<_>>(),
})
})
.collect();
let body = serde_json::to_vec(&serde_json::json!({ "modelSummaries": models }))
.map_err(|e| Error::internal_err(format!("Failed to serialize response: {}", e)))?;
Ok(BedrockProxyResponse {
status_code: StatusCode::OK,
headers: json_response_headers(),
body: BedrockProxyResponseBody::Fixed(Bytes::from(body)),
})
}
async fn list_inference_profiles(
args: &ProxyBuildArgs<'_>,
region: &str,
) -> Result<BedrockProxyResponse, Error> {
let client = create_bedrock_control_client(args, region).await?;
let response =
client.list_inference_profiles().send().await.map_err(|e| {
Error::internal_err(format!("Failed to list inference profiles: {}", e))
})?;
let profiles: Vec<serde_json::Value> = response
.inference_profile_summaries()
.iter()
.map(|p| {
serde_json::json!({
"inferenceProfileId": p.inference_profile_id(),
"inferenceProfileName": p.inference_profile_name(),
"inferenceProfileArn": p.inference_profile_arn(),
"description": p.description(),
"status": p.status().as_str(),
"type": p.r#type().as_str(),
})
})
.collect();
let body = serde_json::to_vec(&serde_json::json!({ "inferenceProfileSummaries": profiles }))
.map_err(|e| Error::internal_err(format!("Failed to serialize response: {}", e)))?;
Ok(BedrockProxyResponse {
status_code: StatusCode::OK,
headers: json_response_headers(),
body: BedrockProxyResponseBody::Fixed(Bytes::from(body)),
})
}
async fn handle_bedrock_sdk_streaming(
model: &str,
body: &[u8],
args: &ProxyBuildArgs<'_>,
region: &str,
) -> Result<BedrockProxyResponse, Error> {
let openai_req: OpenAIRequest = serde_json::from_slice(body)
.map_err(|e| Error::internal_err(format!("Failed to parse OpenAI request: {}", e)))?;
let bedrock_client = create_bedrock_client(args, region).await?;
let enable_prompt_caching = bedrock_model_supports_prompt_caching(model);
let (bedrock_messages, system_prompts) =
openai_messages_to_bedrock(&openai_req.messages, enable_prompt_caching)?;
let inference_config = create_inference_config(openai_req.temperature, openai_req.max_tokens);
let tool_config = build_tool_config_from_request(
openai_req.tools.as_deref(),
openai_req.tool_choice.as_ref(),
enable_prompt_caching,
)?;
let mut request_builder = bedrock_client
.client()
.converse_stream()
.model_id(model)
.set_messages(Some(bedrock_messages));
if !system_prompts.is_empty() {
request_builder = request_builder.set_system(Some(system_prompts));
}
if let Some(config) = inference_config {
request_builder = request_builder.inference_config(config);
}
if let Some(config) = tool_config {
request_builder = request_builder.set_tool_config(Some(config));
}
tracing::debug!("Bedrock SDK streaming: sending converse_stream request");
let stream_output = request_builder.send().await.map_err(|e| {
let error_msg = format!("Bedrock SDK streaming error: {}", format_bedrock_error(&e));
tracing::error!("Bedrock SDK streaming failed: {}", error_msg);
Error::internal_err(error_msg)
})?;
tracing::debug!("Bedrock SDK streaming: stream established successfully");
Ok(BedrockProxyResponse {
status_code: StatusCode::OK,
headers: event_stream_response_headers(),
body: BedrockProxyResponseBody::Stream(
sdk_stream_to_sse(stream_output.stream, model.to_string()).boxed(),
),
})
}
pub fn sdk_stream_to_sse(
stream: aws_sdk_bedrockruntime::primitives::event_stream::EventReceiver<
aws_sdk_bedrockruntime::types::ConverseStreamOutput,
aws_sdk_bedrockruntime::types::error::ConverseStreamOutputError,
>,
model: String,
) -> impl futures::Stream<Item = std::result::Result<Bytes, std::io::Error>> + Send {
let id = format!("chatcmpl-{}", uuid::Uuid::new_v4().simple());
let created = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
struct StreamState {
id: String,
model: String,
created: u64,
tool_calls: HashMap<usize, (String, String, String)>,
current_tool_index: usize,
}
let state = std::sync::Arc::new(tokio::sync::Mutex::new(StreamState {
id,
model,
created,
tool_calls: HashMap::new(),
current_tool_index: 0,
}));
async_stream::stream! {
let mut stream = stream;
let state = state.clone();
loop {
match stream.recv().await {
Ok(Some(event)) => {
let mut state = state.lock().await;
if let Some(tool_call) = bedrock_stream_event_to_tool_start(&event) {
let index = state.current_tool_index;
state.tool_calls.insert(
index,
(tool_call.id.clone(), tool_call.name.clone(), String::new()),
);
let chunk = serde_json::json!({
"id": state.id,
"object": "chat.completion.chunk",
"created": state.created,
"model": state.model,
"choices": [{
"index": 0,
"delta": {
"tool_calls": [{
"index": index,
"id": tool_call.id,
"type": "function",
"function": {
"name": tool_call.name,
"arguments": ""
}
}]
},
"finish_reason": serde_json::Value::Null
}]
});
yield Ok(Bytes::from(format!("data: {}\n\n", chunk)));
}
if let Some(text) = bedrock_stream_event_to_text(&event) {
let chunk = serde_json::json!({
"id": state.id,
"object": "chat.completion.chunk",
"created": state.created,
"model": state.model,
"choices": [{
"index": 0,
"delta": {
"content": text
},
"finish_reason": serde_json::Value::Null
}]
});
yield Ok(Bytes::from(format!("data: {}\n\n", chunk)));
}
if let Some(input_delta) = bedrock_stream_event_to_tool_delta(&event) {
let index = state.current_tool_index;
if let Some((_id, _name, ref mut args)) = state.tool_calls.get_mut(&index) {
args.push_str(&input_delta);
let chunk = serde_json::json!({
"id": state.id,
"object": "chat.completion.chunk",
"created": state.created,
"model": state.model,
"choices": [{
"index": 0,
"delta": {
"tool_calls": [{
"index": index,
"function": {
"arguments": input_delta
}
}]
},
"finish_reason": serde_json::Value::Null
}]
});
yield Ok(Bytes::from(format!("data: {}\n\n", chunk)));
}
}
if bedrock_stream_event_is_block_stop(&event) {
state.current_tool_index += 1;
}
if let aws_sdk_bedrockruntime::types::ConverseStreamOutput::MessageStop(stop) = &event {
let stop_reason = stop.stop_reason().as_str();
let finish_reason = match stop_reason {
"end_turn" => "stop",
"max_tokens" => "length",
"tool_use" => "tool_calls",
"stop_sequence" => "stop",
"guardrail_intervened" | "content_filtered" => "content_filter",
_ => "stop",
};
let chunk = serde_json::json!({
"id": state.id,
"object": "chat.completion.chunk",
"created": state.created,
"model": state.model,
"choices": [{
"index": 0,
"delta": {},
"finish_reason": finish_reason
}]
});
yield Ok(Bytes::from(format!("data: {}\n\n", chunk)));
}
}
Ok(None) => break,
Err(e) => {
yield Err(std::io::Error::new(
std::io::ErrorKind::Other,
e.to_string(),
));
break;
}
}
}
yield Ok(Bytes::from("data: [DONE]\n\n"));
}
}
async fn handle_bedrock_sdk_non_streaming(
model: &str,
body: &[u8],
args: &ProxyBuildArgs<'_>,
region: &str,
) -> Result<BedrockProxyResponse, Error> {
let openai_req: OpenAIRequest = serde_json::from_slice(body)
.map_err(|e| Error::internal_err(format!("Failed to parse OpenAI request: {}", e)))?;
let bedrock_client = create_bedrock_client(args, region).await?;
let enable_prompt_caching = bedrock_model_supports_prompt_caching(model);
let (bedrock_messages, system_prompts) =
openai_messages_to_bedrock(&openai_req.messages, enable_prompt_caching)?;
let inference_config = create_inference_config(openai_req.temperature, openai_req.max_tokens);
let tool_config = build_tool_config_from_request(
openai_req.tools.as_deref(),
openai_req.tool_choice.as_ref(),
enable_prompt_caching,
)?;
let mut request_builder = bedrock_client
.client()
.converse()
.model_id(model)
.set_messages(Some(bedrock_messages));
if !system_prompts.is_empty() {
request_builder = request_builder.set_system(Some(system_prompts));
}
if let Some(config) = inference_config {
request_builder = request_builder.inference_config(config);
}
if let Some(config) = tool_config {
request_builder = request_builder.set_tool_config(Some(config));
}
tracing::debug!("Bedrock SDK non-streaming: sending converse request");
let response = request_builder.send().await.map_err(|e| {
let error_msg = format!(
"Bedrock SDK non-streaming error: {}",
format_bedrock_error(&e)
);
tracing::error!("Bedrock SDK non-streaming failed: {}", error_msg);
Error::internal_err(error_msg)
})?;
tracing::debug!(
"Bedrock SDK non-streaming: response received, stop_reason={}",
response.stop_reason().as_str()
);
let id = format!("chatcmpl-{}", uuid::Uuid::new_v4().simple());
let created = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
let stop_reason = response.stop_reason().as_str();
let finish_reason = match stop_reason {
"end_turn" => "stop",
"max_tokens" => "length",
"tool_use" => "tool_calls",
"stop_sequence" => "stop",
"guardrail_intervened" | "content_filtered" => "content_filter",
_ => "stop",
};
let mut text_content = String::new();
let mut tool_calls: Vec<OpenAIToolCall> = Vec::new();
if let Some(aws_sdk_bedrockruntime::types::ConverseOutput::Message(message)) = response.output()
{
for block in message.content() {
match block {
aws_sdk_bedrockruntime::types::ContentBlock::Text(text) => {
text_content.push_str(text);
}
aws_sdk_bedrockruntime::types::ContentBlock::ToolUse(tool_use) => {
let input_json = document_to_json(tool_use.input());
tool_calls.push(OpenAIToolCall {
id: tool_use.tool_use_id().to_string(),
function: OpenAIFunction {
name: tool_use.name().to_string(),
arguments: serde_json::to_string(&input_json).unwrap_or_default(),
},
r#type: "function".to_string(),
extra_content: None,
});
}
_ => {}
}
}
}
let message = if !tool_calls.is_empty() {
serde_json::json!({
"role": "assistant",
"content": if text_content.is_empty() { serde_json::Value::Null } else { serde_json::Value::String(text_content) },
"tool_calls": tool_calls
})
} else {
serde_json::json!({
"role": "assistant",
"content": text_content
})
};
let usage = if let Some(usage_data) = response.usage() {
serde_json::json!({
"prompt_tokens": usage_data.input_tokens(),
"completion_tokens": usage_data.output_tokens(),
"total_tokens": usage_data.total_tokens()
})
} else {
serde_json::json!({
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 0
})
};
let openai_resp = serde_json::json!({
"id": id,
"object": "chat.completion",
"created": created,
"model": model,
"choices": [{
"index": 0,
"message": message,
"finish_reason": finish_reason
}],
"usage": usage
});
let body = serde_json::to_vec(&openai_resp)
.map_err(|e| Error::internal_err(format!("Failed to serialize OpenAI response: {}", e)))?;
Ok(BedrockProxyResponse {
status_code: StatusCode::OK,
headers: json_response_headers(),
body: BedrockProxyResponseBody::Fixed(Bytes::from(body)),
})
}
fn json_response_headers() -> HeaderMap {
let mut headers = HeaderMap::new();
headers.insert("content-type", "application/json".parse().unwrap());
headers
}
fn event_stream_response_headers() -> HeaderMap {
let mut headers = HeaderMap::new();
headers.insert("content-type", "text/event-stream".parse().unwrap());
headers.insert("cache-control", "no-cache".parse().unwrap());
headers.insert("connection", "keep-alive".parse().unwrap());
headers
}
fn document_to_json(doc: &aws_smithy_types::Document) -> serde_json::Value {
match doc {
aws_smithy_types::Document::Object(map) => {
let mut json_map = serde_json::Map::new();
for (key, value) in map {
json_map.insert(key.clone(), document_to_json(value));
}
serde_json::Value::Object(json_map)
}
aws_smithy_types::Document::Array(values) => {
serde_json::Value::Array(values.iter().map(document_to_json).collect())
}
aws_smithy_types::Document::Number(number) => match number {
aws_smithy_types::Number::PosInt(number) => serde_json::Value::Number((*number).into()),
aws_smithy_types::Number::NegInt(number) => serde_json::Value::Number((*number).into()),
aws_smithy_types::Number::Float(number) => serde_json::json!(*number),
},
aws_smithy_types::Document::String(value) => serde_json::Value::String(value.clone()),
aws_smithy_types::Document::Bool(value) => serde_json::Value::Bool(*value),
aws_smithy_types::Document::Null => serde_json::Value::Null,
}
}
// ============================================================================
// Query Builder
@@ -256,3 +966,60 @@ impl BedrockQueryBuilder {
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn determine_auth_config_prioritizes_bearer_token() {
let config = determine_auth_config(
Some("bearer-token"),
Some("AKIA123"),
Some("secret"),
Some("session-token"),
);
match config {
BedrockAuthConfig::BearerToken(token) => assert_eq!(token, "bearer-token"),
_ => panic!("expected bearer token auth config"),
}
}
#[test]
fn determine_auth_config_uses_iam_with_optional_session_token() {
let config =
determine_auth_config(None, Some("AKIA123"), Some("secret"), Some("session-token"));
match config {
BedrockAuthConfig::IamCredentials {
access_key_id,
secret_access_key,
session_token,
} => {
assert_eq!(access_key_id, "AKIA123");
assert_eq!(secret_access_key, "secret");
assert_eq!(session_token.as_deref(), Some("session-token"));
}
_ => panic!("expected IAM auth config"),
}
}
#[test]
fn determine_auth_config_treats_empty_session_token_as_none() {
let config = determine_auth_config(None, Some("AKIA123"), Some("secret"), Some(""));
match config {
BedrockAuthConfig::IamCredentials { session_token, .. } => {
assert!(session_token.is_none());
}
_ => panic!("expected IAM auth config"),
}
}
#[test]
fn determine_auth_config_falls_back_to_environment() {
let config = determine_auth_config(None, Some("AKIA123"), None, Some("session-token"));
assert!(matches!(config, BedrockAuthConfig::Environment));
}
}
@@ -408,6 +408,8 @@ fn build_google_ai_model_endpoint(
action: &str,
is_vertex: bool,
) -> String {
let model = model.strip_prefix("models/").unwrap_or(model);
if is_vertex {
format!("{}/{}:{}", base_url, model, action)
} else {
@@ -416,6 +418,10 @@ fn build_google_ai_model_endpoint(
}
fn add_google_ai_auth_header(headers: &mut Vec<(String, String)>, api_key: &str, is_vertex: bool) {
// Native Google AI proxy intentionally does not apply AI_HTTP_HEADERS or
// resource custom headers yet. Gemini/Vertex header semantics are
// provider-specific; keep this limited to required auth headers until
// explicit custom-header support is designed.
if is_vertex {
headers.push(("Authorization".to_string(), format!("Bearer {}", api_key)));
} else {
@@ -749,6 +755,32 @@ mod tests {
assert!(body["contents"].is_array());
}
#[test]
fn builds_standard_google_ai_endpoint_from_model_resource_name() {
assert_eq!(
build_google_ai_model_endpoint(
"https://generativelanguage.googleapis.com/v1beta",
"models/gemini-2.0-flash",
"generateContent",
false,
),
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.0-flash:generateContent"
);
}
#[test]
fn builds_vertex_google_ai_endpoint_from_model_resource_name() {
assert_eq!(
build_google_ai_model_endpoint(
"https://us-central1-aiplatform.googleapis.com/v1/projects/p/locations/us-central1/publishers/google/models",
"models/gemini-2.0-flash",
"streamGenerateContent",
true,
),
"https://us-central1-aiplatform.googleapis.com/v1/projects/p/locations/us-central1/publishers/google/models/gemini-2.0-flash:streamGenerateContent"
);
}
#[test]
fn builds_vertex_google_ai_streaming_proxy_request() {
let credentials = credentials(
+1 -6
View File
@@ -40,7 +40,7 @@ gcp_trigger = ["dep:windmill-trigger-gcp", "windmill-store/gcp_trigger"]
azure_trigger = ["dep:windmill-trigger-azure", "windmill-store/azure_trigger"]
cloud = ["windmill-common/cloud", "windmill-api-auth/cloud", "windmill-store/cloud", "windmill-api-workspaces/cloud"]
mcp = ["dep:windmill-mcp", "windmill-mcp/server", "windmill-mcp/auth", "windmill-api-auth/mcp", "windmill-store/mcp"]
bedrock = ["windmill-ai/bedrock", "dep:aws-sdk-bedrock", "dep:aws-sdk-bedrockruntime", "dep:aws-config", "dep:aws-credential-types", "dep:aws-smithy-types"]
bedrock = ["windmill-ai/bedrock"]
python = ["windmill-dep-map/python", "dep:windmill-parser-py", "dep:windmill-parser-py-imports", "windmill-api-scripts/python", "windmill-api-configs/python", "windmill-api-agent-workers?/python", "windmill-trigger/python", "windmill-common/python"]
no_auth = ["windmill-api-auth/no_auth", "windmill-store/no_auth", "windmill-api-users/no_auth"]
quickjs = ["windmill-jseval/quickjs"]
@@ -172,11 +172,6 @@ rustls = { workspace = true }
aws-sigv4 = { workspace = true, optional = true }
aws-sdk-config = { workspace = true, optional = true }
aws-config = { workspace = true, optional = true }
aws-credential-types = { workspace = true, optional = true }
aws-sdk-bedrock = { workspace = true, optional = true }
aws-sdk-bedrockruntime = { workspace = true, optional = true }
aws-smithy-types = { workspace = true, optional = true }
async-trait.workspace = true
eventsource-stream.workspace = true
windmill-jseval.workspace = true
+39 -89
View File
@@ -1,5 +1,3 @@
#[cfg(feature = "bedrock")]
use crate::bedrock;
use crate::db::{ApiAuthed, DB};
use crate::utils::check_scopes;
@@ -20,6 +18,10 @@ use windmill_ai::ai_cache::current_instance_ai_config_revision;
use windmill_ai::ai_providers::{
empty_string_as_none, AIPlatform, AIProvider, ProviderConfig, ProviderModel,
};
#[cfg(feature = "bedrock")]
use windmill_ai::providers::bedrock::{
handle_bedrock_proxy, BedrockProxyResponse, BedrockProxyResponseBody,
};
use windmill_ai::providers::{
create_proxy_query_builder,
google_ai::{
@@ -540,6 +542,18 @@ fn google_ai_proxy_response_to_body(
(response.status_code, response.headers, body)
}
#[cfg(feature = "bedrock")]
fn bedrock_proxy_response_to_body(
response: BedrockProxyResponse,
) -> (http::StatusCode, HeaderMap, axum::body::Body) {
let body = match response.body {
BedrockProxyResponseBody::Fixed(body) => axum::body::Body::from(body),
BedrockProxyResponseBody::Stream(stream) => axum::body::Body::from_stream(stream),
};
(response.status_code, response.headers, body)
}
pub(crate) fn inject_keepalives<S>(
upstream: S,
interval: Duration,
@@ -885,95 +899,31 @@ async fn proxy(
// Handle Bedrock-specific logic when the feature is enabled
#[cfg(feature = "bedrock")]
{
// Extract model and streaming flag for Bedrock transformation (only for POST requests)
let (model, is_streaming) = if matches!(proxy_mode, ProxyExecutionMode::NativeAwsBedrock)
&& method == Method::POST
{
#[derive(Deserialize, Debug)]
struct BedrockRequest {
model: String,
#[serde(default)]
stream: bool,
}
let parsed: BedrockRequest = serde_json::from_slice(&body)
.map_err(|e| Error::internal_err(format!("Failed to parse request body: {}", e)))?;
(Some(parsed.model), parsed.stream)
} else {
(None, false)
};
if matches!(proxy_mode, ProxyExecutionMode::NativeAwsBedrock) {
let mut tx = db.begin().await?;
audit_log(
&mut *tx,
&authed,
"ai.request",
ActionKind::Execute,
&w_id,
Some(&authed.email),
Some([("ai_config_path", &format!("{:?}", ai_path)[..])].into()),
)
.await?;
tx.commit().await?;
// For Bedrock requests, use the SDK-based approach
if matches!(proxy_mode, ProxyExecutionMode::NativeAwsBedrock) {
let region = request_config
.region
.as_deref()
.unwrap_or(windmill_ai::ai_providers::USE_ENV_REGION);
let credentials = request_config.into_provider_credentials(provider.clone());
let response = handle_bedrock_proxy(&ProxyBuildArgs {
method: &method,
path: &ai_path,
headers: &headers,
body: &body,
credentials: &credentials,
})
.await?;
// Audit log before making the SDK request
let mut tx = db.begin().await?;
audit_log(
&mut *tx,
&authed,
"ai.request",
ActionKind::Execute,
&w_id,
Some(&authed.email),
Some([("ai_config_path", &format!("{:?}", ai_path)[..])].into()),
)
.await?;
tx.commit().await?;
// Handle GET requests for control plane operations
if method == Method::GET {
if ai_path == "foundation-models" {
return bedrock::list_foundation_models(
request_config.api_key.as_deref(),
request_config.aws_access_key_id.as_deref(),
request_config.aws_secret_access_key.as_deref(),
request_config.aws_session_token.as_deref(),
region,
)
.await;
} else if ai_path == "inference-profiles" {
return bedrock::list_inference_profiles(
request_config.api_key.as_deref(),
request_config.aws_access_key_id.as_deref(),
request_config.aws_secret_access_key.as_deref(),
request_config.aws_session_token.as_deref(),
region,
)
.await;
}
}
// Handle POST requests for inference
if method == Method::POST && model.is_some() {
if is_streaming {
return bedrock::handle_bedrock_sdk_streaming(
model.as_ref().unwrap(),
&body,
request_config.api_key.as_deref(),
request_config.aws_access_key_id.as_deref(),
request_config.aws_secret_access_key.as_deref(),
request_config.aws_session_token.as_deref(),
region,
)
.await;
} else {
return bedrock::handle_bedrock_sdk_non_streaming(
model.as_ref().unwrap(),
&body,
request_config.api_key.as_deref(),
request_config.aws_access_key_id.as_deref(),
request_config.aws_secret_access_key.as_deref(),
request_config.aws_session_token.as_deref(),
region,
)
.await;
}
}
}
return Ok(bedrock_proxy_response_to_body(response));
}
// When bedrock feature is disabled, return error for Bedrock provider
-873
View File
@@ -1,873 +0,0 @@
//! AWS Bedrock SDK-based operations for the AI chat proxy.
//!
//! This module provides SDK-based request handling for Bedrock:
//!
//! ## Inference (Runtime SDK):
//! - `handle_bedrock_sdk_streaming`: Uses BedrockClient for streaming requests
//! - `handle_bedrock_sdk_non_streaming`: Uses BedrockClient for non-streaming requests
//! - `sdk_stream_to_sse`: Converts SDK ConverseStream events to SSE format
//!
//! ## Control Plane (Bedrock SDK):
//! - `list_foundation_models`: Lists available foundation models
//! - `list_inference_profiles`: Lists inference profiles
//!
//! Shared AWS SDK code is available in `windmill_common::ai_bedrock`, including:
//! - `BedrockClient`: SDK wrapper with bearer token and IAM auth
//! - Stream event parsing functions
//! - Helper utilities
use axum::body::Bytes;
use serde::Deserialize;
use windmill_ai::ai_bedrock::build_tool_config;
use windmill_ai::ai_bedrock::{
bedrock_stream_event_is_block_stop, bedrock_stream_event_to_text,
bedrock_stream_event_to_tool_delta, bedrock_stream_event_to_tool_start, format_bedrock_error,
BedrockClient,
};
use windmill_ai::ai_types::{
OpenAIFunction, OpenAIMessage, OpenAIToolCall, ToolDef, ToolDefFunction,
};
use windmill_common::error::{Error, Result};
// ============================================================================
// Shared Request Types for SDK-Based Handlers
// ============================================================================
/// OpenAI-format request body for Bedrock SDK handlers
#[derive(Deserialize, Debug)]
struct OpenAIRequest {
messages: Vec<OpenAIMessage>,
#[serde(default)]
tools: Option<Vec<OpenAIToolDef>>,
#[serde(default)]
tool_choice: Option<serde_json::Value>,
#[serde(default)]
max_tokens: Option<i32>,
#[serde(default)]
temperature: Option<f32>,
}
#[derive(Deserialize, Debug)]
struct OpenAIToolDef {
#[serde(default)]
#[allow(dead_code)]
r#type: Option<String>,
function: OpenAIToolFunction,
}
#[derive(Deserialize, Debug)]
struct OpenAIToolFunction {
name: String,
#[serde(default)]
description: Option<String>,
#[serde(default)]
parameters: Option<serde_json::Value>,
}
// ============================================================================
// Shared Helper Functions for SDK-Based Handlers
// ============================================================================
/// Authentication configuration for Bedrock clients
enum BedrockAuthConfig {
BearerToken(String),
IamCredentials {
access_key_id: String,
secret_access_key: String,
session_token: Option<String>,
},
Environment,
}
/// Determine auth configuration with priority: bearer token → IAM credentials → environment
fn determine_auth_config(
api_key: Option<&str>,
aws_access_key_id: Option<&str>,
aws_secret_access_key: Option<&str>,
aws_session_token: Option<&str>,
) -> BedrockAuthConfig {
if let Some(key) = api_key.filter(|k| !k.is_empty()) {
BedrockAuthConfig::BearerToken(key.to_string())
} else if let (Some(access_key_id), Some(secret_access_key)) = (
aws_access_key_id.filter(|s| !s.is_empty()),
aws_secret_access_key.filter(|s| !s.is_empty()),
) {
BedrockAuthConfig::IamCredentials {
access_key_id: access_key_id.to_string(),
secret_access_key: secret_access_key.to_string(),
session_token: aws_session_token
.filter(|token| !token.is_empty())
.map(str::to_string),
}
} else {
BedrockAuthConfig::Environment
}
}
/// Create a BedrockClient with auth priority: bearer token → IAM credentials → environment
async fn create_bedrock_client(
api_key: Option<&str>,
aws_access_key_id: Option<&str>,
aws_secret_access_key: Option<&str>,
aws_session_token: Option<&str>,
region: &str,
) -> Result<BedrockClient> {
match determine_auth_config(
api_key,
aws_access_key_id,
aws_secret_access_key,
aws_session_token,
) {
BedrockAuthConfig::BearerToken(key) => BedrockClient::from_bearer_token(key, region).await,
BedrockAuthConfig::IamCredentials { access_key_id, secret_access_key, session_token } => {
BedrockClient::from_credentials(access_key_id, secret_access_key, session_token, region)
.await
}
BedrockAuthConfig::Environment => BedrockClient::from_env(region).await,
}
}
/// Convert OpenAIToolDef array to tool configuration for Bedrock SDK
fn build_tool_config_from_request(
tools: Option<&[OpenAIToolDef]>,
tool_choice: Option<&serde_json::Value>,
enable_prompt_caching: bool,
) -> Result<Option<aws_sdk_bedrockruntime::types::ToolConfiguration>> {
if let Some(tools) = tools {
let tool_defs: Vec<ToolDef> = tools
.iter()
.map(|t| ToolDef {
r#type: "function".to_string(),
function: ToolDefFunction {
name: t.function.name.clone(),
description: t.function.description.clone(),
parameters: Box::from(
serde_json::value::RawValue::from_string(
serde_json::to_string(
&t.function
.parameters
.clone()
.unwrap_or(serde_json::json!({})),
)
.unwrap_or_default(),
)
.unwrap_or_else(|_| {
serde_json::value::RawValue::from_string("{}".to_string()).unwrap()
}),
),
},
})
.collect();
// Determine if we should force tool use based on tool_choice
let force_tool_use = tool_choice
.map(|tc| tc == "required" || tc.as_str() == Some("required"))
.unwrap_or(false);
build_tool_config(Some(&tool_defs), force_tool_use, enable_prompt_caching)
} else {
Ok(None)
}
}
// ============================================================================
// Control Plane Operations (using aws-sdk-bedrock)
// ============================================================================
/// Create a Bedrock control plane client with auth priority: bearer token → IAM credentials → environment
async fn create_bedrock_control_client(
api_key: Option<&str>,
aws_access_key_id: Option<&str>,
aws_secret_access_key: Option<&str>,
aws_session_token: Option<&str>,
region: &str,
) -> Result<aws_sdk_bedrock::Client> {
use aws_config::BehaviorVersion;
use windmill_ai::ai_bedrock::BearerTokenProvider;
let region_provider = aws_sdk_bedrock::config::Region::new(region.to_string());
match determine_auth_config(
api_key,
aws_access_key_id,
aws_secret_access_key,
aws_session_token,
) {
BedrockAuthConfig::BearerToken(key) => {
let config = aws_sdk_bedrock::config::Builder::new()
.region(region_provider)
.behavior_version(BehaviorVersion::latest())
.token_provider(BearerTokenProvider::new(key))
.build();
Ok(aws_sdk_bedrock::Client::from_conf(config))
}
BedrockAuthConfig::IamCredentials { access_key_id, secret_access_key, session_token } => {
let credentials = aws_credential_types::Credentials::new(
access_key_id,
secret_access_key,
session_token,
None,
"windmill",
);
let config = aws_sdk_bedrock::config::Builder::new()
.region(region_provider)
.behavior_version(BehaviorVersion::latest())
.credentials_provider(credentials)
.build();
Ok(aws_sdk_bedrock::Client::from_conf(config))
}
BedrockAuthConfig::Environment => {
let config = aws_config::defaults(BehaviorVersion::latest())
.region(region_provider)
.load()
.await;
Ok(aws_sdk_bedrock::Client::new(&config))
}
}
}
/// List foundation models using the Bedrock SDK
pub async fn list_foundation_models(
api_key: Option<&str>,
aws_access_key_id: Option<&str>,
aws_secret_access_key: Option<&str>,
aws_session_token: Option<&str>,
region: &str,
) -> Result<(http::StatusCode, http::HeaderMap, axum::body::Body)> {
let client = create_bedrock_control_client(
api_key,
aws_access_key_id,
aws_secret_access_key,
aws_session_token,
region,
)
.await?;
let response = client
.list_foundation_models()
.send()
.await
.map_err(|e| Error::internal_err(format!("Failed to list foundation models: {}", e)))?;
// Convert to JSON response
let models: Vec<serde_json::Value> = response
.model_summaries()
.iter()
.map(|m| {
serde_json::json!({
"modelId": m.model_id(),
"modelName": m.model_name(),
"providerName": m.provider_name(),
"modelArn": m.model_arn(),
"inputModalities": m.input_modalities().iter().map(|i| i.as_str()).collect::<Vec<_>>(),
"outputModalities": m.output_modalities().iter().map(|o| o.as_str()).collect::<Vec<_>>(),
"responseStreamingSupported": m.response_streaming_supported(),
"inferenceTypesSupported": m.inference_types_supported().iter().map(|i| i.as_str()).collect::<Vec<_>>(),
})
})
.collect();
let body = serde_json::json!({ "modelSummaries": models });
let body_bytes = serde_json::to_vec(&body)
.map_err(|e| Error::internal_err(format!("Failed to serialize response: {}", e)))?;
let mut headers = http::HeaderMap::new();
headers.insert("content-type", "application/json".parse().unwrap());
Ok((
http::StatusCode::OK,
headers,
axum::body::Body::from(body_bytes),
))
}
/// List inference profiles using the Bedrock SDK
pub async fn list_inference_profiles(
api_key: Option<&str>,
aws_access_key_id: Option<&str>,
aws_secret_access_key: Option<&str>,
aws_session_token: Option<&str>,
region: &str,
) -> Result<(http::StatusCode, http::HeaderMap, axum::body::Body)> {
let client = create_bedrock_control_client(
api_key,
aws_access_key_id,
aws_secret_access_key,
aws_session_token,
region,
)
.await?;
let response =
client.list_inference_profiles().send().await.map_err(|e| {
Error::internal_err(format!("Failed to list inference profiles: {}", e))
})?;
// Convert to JSON response
let profiles: Vec<serde_json::Value> = response
.inference_profile_summaries()
.iter()
.map(|p| {
serde_json::json!({
"inferenceProfileId": p.inference_profile_id(),
"inferenceProfileName": p.inference_profile_name(),
"inferenceProfileArn": p.inference_profile_arn(),
"description": p.description(),
"status": p.status().as_str(),
"type": p.r#type().as_str(),
})
})
.collect();
let body = serde_json::json!({ "inferenceProfileSummaries": profiles });
let body_bytes = serde_json::to_vec(&body)
.map_err(|e| Error::internal_err(format!("Failed to serialize response: {}", e)))?;
let mut headers = http::HeaderMap::new();
headers.insert("content-type", "application/json".parse().unwrap());
Ok((
http::StatusCode::OK,
headers,
axum::body::Body::from(body_bytes),
))
}
// ============================================================================
// Inference Operations (using aws-sdk-bedrockruntime)
// ============================================================================
/// Handle Bedrock streaming request using the AWS SDK.
///
/// This function uses the shared BedrockClient to make streaming requests
/// and converts the SDK stream events to SSE format for the proxy response.
///
/// Auth priority: bearer token → IAM credentials → environment credentials
pub async fn handle_bedrock_sdk_streaming(
model: &str,
body: &Bytes,
api_key: Option<&str>,
aws_access_key_id: Option<&str>,
aws_secret_access_key: Option<&str>,
aws_session_token: Option<&str>,
region: &str,
) -> Result<(http::StatusCode, http::HeaderMap, axum::body::Body)> {
let openai_req: OpenAIRequest = serde_json::from_slice(body)
.map_err(|e| Error::internal_err(format!("Failed to parse OpenAI request: {}", e)))?;
// Create Bedrock client using shared helper
let bedrock_client = create_bedrock_client(
api_key,
aws_access_key_id,
aws_secret_access_key,
aws_session_token,
region,
)
.await?;
// Convert messages using shared conversion
let enable_prompt_caching =
windmill_ai::ai_bedrock::bedrock_model_supports_prompt_caching(model);
let (bedrock_messages, system_prompts) =
windmill_ai::ai_bedrock::openai_messages_to_bedrock(
&openai_req.messages,
enable_prompt_caching,
)?;
// Build inference configuration
let inference_config = windmill_ai::ai_bedrock::create_inference_config(
openai_req.temperature,
openai_req.max_tokens,
);
// Convert tools using shared helper
let tool_config = build_tool_config_from_request(
openai_req.tools.as_deref(),
openai_req.tool_choice.as_ref(),
enable_prompt_caching,
)?;
// Build the SDK request
let mut request_builder = bedrock_client
.client()
.converse_stream()
.model_id(model)
.set_messages(Some(bedrock_messages));
if !system_prompts.is_empty() {
request_builder = request_builder.set_system(Some(system_prompts));
}
if let Some(config) = inference_config {
request_builder = request_builder.inference_config(config);
}
if let Some(config) = tool_config {
request_builder = request_builder.set_tool_config(Some(config));
}
// Send the request and get the stream
tracing::debug!("Bedrock SDK streaming: sending converse_stream request");
let stream_output = request_builder.send().await.map_err(|e| {
let error_msg = format!("Bedrock SDK streaming error: {}", format_bedrock_error(&e));
tracing::error!("Bedrock SDK streaming failed: {}", error_msg);
Error::internal_err(error_msg)
})?;
tracing::debug!("Bedrock SDK streaming: stream established successfully");
// Convert SDK stream to SSE (pass the inner stream, not the full output)
let sse_stream = sdk_stream_to_sse(stream_output.stream, model.to_string());
// Build response headers
let mut response_headers = http::HeaderMap::new();
response_headers.insert("content-type", "text/event-stream".parse().unwrap());
response_headers.insert("cache-control", "no-cache".parse().unwrap());
response_headers.insert("connection", "keep-alive".parse().unwrap());
Ok((
http::StatusCode::OK,
response_headers,
axum::body::Body::from_stream(sse_stream),
))
}
/// Convert AWS SDK ConverseStream events to SSE format.
///
/// Uses shared stream parsing functions from windmill_common::ai_bedrock
/// to extract text deltas and tool calls from the SDK stream events.
pub fn sdk_stream_to_sse(
stream: aws_sdk_bedrockruntime::primitives::event_stream::EventReceiver<
aws_sdk_bedrockruntime::types::ConverseStreamOutput,
aws_sdk_bedrockruntime::types::error::ConverseStreamOutputError,
>,
model: String,
) -> impl futures::Stream<Item = std::result::Result<bytes::Bytes, std::io::Error>> + Send {
use std::collections::HashMap;
let id = format!("chatcmpl-{}", uuid::Uuid::new_v4().simple());
let created = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
// State to track partial tool calls
struct StreamState {
id: String,
model: String,
created: u64,
tool_calls: HashMap<usize, (String, String, String)>, // index -> (id, name, args)
current_tool_index: usize,
}
let state = std::sync::Arc::new(tokio::sync::Mutex::new(StreamState {
id: id.clone(),
model: model.clone(),
created,
tool_calls: HashMap::new(),
current_tool_index: 0,
}));
async_stream::stream! {
let mut stream = stream;
let state = state.clone();
loop {
match stream.recv().await {
Ok(Some(event)) => {
let mut state = state.lock().await;
// Handle tool use start
if let Some(tool_call) = bedrock_stream_event_to_tool_start(&event) {
let index = state.current_tool_index;
state.tool_calls.insert(
index,
(tool_call.id.clone(), tool_call.name.clone(), String::new()),
);
// Send initial tool call chunk
let chunk = serde_json::json!({
"id": state.id,
"object": "chat.completion.chunk",
"created": state.created,
"model": state.model,
"choices": [{
"index": 0,
"delta": {
"tool_calls": [{
"index": index,
"id": tool_call.id,
"type": "function",
"function": {
"name": tool_call.name,
"arguments": ""
}
}]
},
"finish_reason": serde_json::Value::Null
}]
});
yield Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk)));
}
// Handle text delta
if let Some(text) = bedrock_stream_event_to_text(&event) {
let chunk = serde_json::json!({
"id": state.id,
"object": "chat.completion.chunk",
"created": state.created,
"model": state.model,
"choices": [{
"index": 0,
"delta": {
"content": text
},
"finish_reason": serde_json::Value::Null
}]
});
yield Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk)));
}
// Handle tool use input delta
if let Some(input_delta) = bedrock_stream_event_to_tool_delta(&event) {
let index = state.current_tool_index;
if let Some((_id, _name, ref mut args)) = state.tool_calls.get_mut(&index) {
args.push_str(&input_delta);
let chunk = serde_json::json!({
"id": state.id,
"object": "chat.completion.chunk",
"created": state.created,
"model": state.model,
"choices": [{
"index": 0,
"delta": {
"tool_calls": [{
"index": index,
"function": {
"arguments": input_delta
}
}]
},
"finish_reason": serde_json::Value::Null
}]
});
yield Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk)));
}
}
// Handle content block stop
if bedrock_stream_event_is_block_stop(&event) {
state.current_tool_index += 1;
}
// Handle message stop
if let aws_sdk_bedrockruntime::types::ConverseStreamOutput::MessageStop(stop) = &event {
let stop_reason = stop.stop_reason().as_str();
let finish_reason = match stop_reason {
"end_turn" => "stop",
"max_tokens" => "length",
"tool_use" => "tool_calls",
"stop_sequence" => "stop",
"guardrail_intervened" | "content_filtered" => "content_filter",
_ => "stop",
};
let chunk = serde_json::json!({
"id": state.id,
"object": "chat.completion.chunk",
"created": state.created,
"model": state.model,
"choices": [{
"index": 0,
"delta": {},
"finish_reason": finish_reason
}]
});
yield Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk)));
}
}
Ok(None) => break,
Err(e) => {
yield Err(std::io::Error::new(
std::io::ErrorKind::Other,
e.to_string(),
));
break;
}
}
}
// Send [DONE] at the end
yield Ok(bytes::Bytes::from("data: [DONE]\n\n"));
}
}
/// Handle non-streaming Bedrock request using the AWS SDK.
///
/// Auth priority: bearer token → IAM credentials → environment credentials
pub async fn handle_bedrock_sdk_non_streaming(
model: &str,
body: &Bytes,
api_key: Option<&str>,
aws_access_key_id: Option<&str>,
aws_secret_access_key: Option<&str>,
aws_session_token: Option<&str>,
region: &str,
) -> Result<(http::StatusCode, http::HeaderMap, axum::body::Body)> {
let openai_req: OpenAIRequest = serde_json::from_slice(body)
.map_err(|e| Error::internal_err(format!("Failed to parse OpenAI request: {}", e)))?;
// Create Bedrock client using shared helper
let bedrock_client = create_bedrock_client(
api_key,
aws_access_key_id,
aws_secret_access_key,
aws_session_token,
region,
)
.await?;
// Convert messages using shared conversion
let enable_prompt_caching =
windmill_ai::ai_bedrock::bedrock_model_supports_prompt_caching(model);
let (bedrock_messages, system_prompts) =
windmill_ai::ai_bedrock::openai_messages_to_bedrock(
&openai_req.messages,
enable_prompt_caching,
)?;
// Build inference configuration
let inference_config = windmill_ai::ai_bedrock::create_inference_config(
openai_req.temperature,
openai_req.max_tokens,
);
// Convert tools using shared helper
let tool_config = build_tool_config_from_request(
openai_req.tools.as_deref(),
openai_req.tool_choice.as_ref(),
enable_prompt_caching,
)?;
// Build the SDK request (non-streaming)
let mut request_builder = bedrock_client
.client()
.converse()
.model_id(model)
.set_messages(Some(bedrock_messages));
if !system_prompts.is_empty() {
request_builder = request_builder.set_system(Some(system_prompts));
}
if let Some(config) = inference_config {
request_builder = request_builder.inference_config(config);
}
if let Some(config) = tool_config {
request_builder = request_builder.set_tool_config(Some(config));
}
// Send the request
tracing::debug!("Bedrock SDK non-streaming: sending converse request");
let response = request_builder.send().await.map_err(|e| {
let error_msg = format!(
"Bedrock SDK non-streaming error: {}",
format_bedrock_error(&e)
);
tracing::error!("Bedrock SDK non-streaming failed: {}", error_msg);
Error::internal_err(error_msg)
})?;
tracing::debug!(
"Bedrock SDK non-streaming: response received, stop_reason={}",
response.stop_reason().as_str()
);
// Convert response to OpenAI format
let id = format!("chatcmpl-{}", uuid::Uuid::new_v4().simple());
let created = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
// Extract stop reason
let stop_reason = response.stop_reason().as_str();
let finish_reason = match stop_reason {
"end_turn" => "stop",
"max_tokens" => "length",
"tool_use" => "tool_calls",
"stop_sequence" => "stop",
"guardrail_intervened" | "content_filtered" => "content_filter",
_ => "stop",
};
// Extract message content
let mut text_content = String::new();
let mut tool_calls: Vec<OpenAIToolCall> = Vec::new();
if let Some(output) = response.output() {
if let aws_sdk_bedrockruntime::types::ConverseOutput::Message(message) = output {
for block in message.content() {
match block {
aws_sdk_bedrockruntime::types::ContentBlock::Text(text) => {
text_content.push_str(text);
}
aws_sdk_bedrockruntime::types::ContentBlock::ToolUse(tool_use) => {
// Convert Document back to JSON string
let input_json = document_to_json(tool_use.input());
tool_calls.push(OpenAIToolCall {
id: tool_use.tool_use_id().to_string(),
function: OpenAIFunction {
name: tool_use.name().to_string(),
arguments: serde_json::to_string(&input_json).unwrap_or_default(),
},
r#type: "function".to_string(),
extra_content: None,
});
}
_ => {}
}
}
}
}
// Build the message
let message = if !tool_calls.is_empty() {
serde_json::json!({
"role": "assistant",
"content": if text_content.is_empty() { serde_json::Value::Null } else { serde_json::Value::String(text_content) },
"tool_calls": tool_calls
})
} else {
serde_json::json!({
"role": "assistant",
"content": text_content
})
};
// Extract usage information
let usage = if let Some(usage_data) = response.usage() {
serde_json::json!({
"prompt_tokens": usage_data.input_tokens(),
"completion_tokens": usage_data.output_tokens(),
"total_tokens": usage_data.total_tokens()
})
} else {
serde_json::json!({
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 0
})
};
// Build OpenAI-format response
let openai_resp = serde_json::json!({
"id": id,
"object": "chat.completion",
"created": created,
"model": model,
"choices": [{
"index": 0,
"message": message,
"finish_reason": finish_reason
}],
"usage": usage
});
let response_body = serde_json::to_vec(&openai_resp)
.map_err(|e| Error::internal_err(format!("Failed to serialize OpenAI response: {}", e)))?;
let mut response_headers = http::HeaderMap::new();
response_headers.insert("content-type", "application/json".parse().unwrap());
Ok((
http::StatusCode::OK,
response_headers,
axum::body::Body::from(response_body),
))
}
/// Convert AWS Smithy Document to serde_json::Value
fn document_to_json(doc: &aws_smithy_types::Document) -> serde_json::Value {
match doc {
aws_smithy_types::Document::Object(map) => {
let mut json_map = serde_json::Map::new();
for (k, v) in map {
json_map.insert(k.clone(), document_to_json(v));
}
serde_json::Value::Object(json_map)
}
aws_smithy_types::Document::Array(arr) => {
serde_json::Value::Array(arr.iter().map(document_to_json).collect())
}
aws_smithy_types::Document::Number(num) => match num {
aws_smithy_types::Number::PosInt(n) => serde_json::Value::Number((*n).into()),
aws_smithy_types::Number::NegInt(n) => serde_json::Value::Number((*n).into()),
aws_smithy_types::Number::Float(f) => serde_json::json!(*f),
},
aws_smithy_types::Document::String(s) => serde_json::Value::String(s.clone()),
aws_smithy_types::Document::Bool(b) => serde_json::Value::Bool(*b),
aws_smithy_types::Document::Null => serde_json::Value::Null,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn determine_auth_config_prioritizes_bearer_token() {
let config = determine_auth_config(
Some("bearer-token"),
Some("AKIA123"),
Some("secret"),
Some("session-token"),
);
match config {
BedrockAuthConfig::BearerToken(token) => assert_eq!(token, "bearer-token"),
_ => panic!("expected bearer token auth config"),
}
}
#[test]
fn determine_auth_config_uses_iam_with_optional_session_token() {
let config =
determine_auth_config(None, Some("AKIA123"), Some("secret"), Some("session-token"));
match config {
BedrockAuthConfig::IamCredentials {
access_key_id,
secret_access_key,
session_token,
} => {
assert_eq!(access_key_id, "AKIA123");
assert_eq!(secret_access_key, "secret");
assert_eq!(session_token.as_deref(), Some("session-token"));
}
_ => panic!("expected IAM auth config"),
}
}
#[test]
fn determine_auth_config_treats_empty_session_token_as_none() {
let config = determine_auth_config(None, Some("AKIA123"), Some("secret"), Some(""));
match config {
BedrockAuthConfig::IamCredentials { session_token, .. } => {
assert!(session_token.is_none());
}
_ => panic!("expected IAM auth config"),
}
}
#[test]
fn determine_auth_config_falls_back_to_environment() {
let config = determine_auth_config(None, Some("AKIA123"), None, Some("session-token"));
assert!(matches!(config, BedrockAuthConfig::Environment));
}
}
-2
View File
@@ -73,8 +73,6 @@ pub mod auth;
#[cfg(all(feature = "private", feature = "parquet"))]
pub mod azure_proxy_ee;
mod azure_proxy_oss;
#[cfg(feature = "bedrock")]
mod bedrock;
mod capture;
mod concurrency_groups;
mod db;
+73 -14
View File
@@ -5,10 +5,10 @@
AI provider logic is currently split across three crates with duplicate code:
- **windmill-common** — base types (`ai_types`, `ai_providers`, `ai_google`, `ai_bedrock`, `ai_cache`)
- **windmill-api** — chat proxy (`ai.rs`, `google.rs`, `bedrock.rs`) with its own request building for Google/Bedrock, plus `AIRequestConfig::prepare_request` for auth/URL handling
- **windmill-api** — chat proxy routes (`ai.rs`), audit logging, caching, and DB-backed credential resolution through `AIRequestConfig`
- **windmill-worker** — agent execution (`ai/` module) with `QueryBuilder` trait, SSE parsers, provider implementations
The goal: a single `windmill-ai` crate with all AI provider logic. Both the API proxy and worker agent use `QueryBuilder` for every provider — no more duplicate logic.
The goal: a single `windmill-ai` crate with all AI provider logic. Worker agent execution uses `QueryBuilder`; the API proxy uses `QueryBuilder::build_proxy_request` for HTTP-forwarding providers and native proxy handlers for providers that need response conversion or SDK execution.
## Dependency Direction
@@ -25,7 +25,7 @@ windmill-common does **NOT** re-export from windmill-ai (would be circular). All
## Reviewer Note: Keep API Proxy Unification Split
The crate boundary, shared utilities, SSE parsers, image handling, and worker provider implementations are now in `windmill-ai`. The remaining duplication is the API proxy path: `AIRequestConfig::prepare_request`, `windmill-api/src/google.rs`, and `windmill-api/src/bedrock.rs` still own API-specific request transformation.
The crate boundary, shared utilities, SSE parsers, image handling, worker provider implementations, and provider-specific API proxy transformations are now in `windmill-ai`. The remaining duplication is credential shape and resolution: `windmill-api` still resolves DB-backed proxy credentials through `AIRequestConfig`, while worker agent execution still receives `ProviderWithResource`.
Do not jump directly from the current state to full proxy and credential unification in one PR. The API proxy combines request transformation, endpoint selection, auth headers, custom headers, OAuth user injection, Azure URL handling, Anthropic Vertex handling, Bedrock SDK calls, and SSE keepalive behavior. Split the work by risk:
- Introduce shared proxy request and credential types first.
@@ -70,7 +70,7 @@ Follow-up status: Anthropic/Vertex proxy handling has since moved into
`windmill-ai`, and the dead `AIRequestConfig::prepare_request` fallback has
been removed.
## Current Phase PR: Proxy Execution Mode + Google AI Proxy Migration
## Completed Phase: Proxy Execution Mode + Google AI Proxy Migration
Goal: introduce a shared provider execution classifier before moving Google AI
and Bedrock. `ProxyRequest` is a good contract for HTTP-forwarding providers
@@ -101,6 +101,62 @@ Validation:
- `cargo test -p windmill-api maps_request_config_to_provider_credentials`
- `cargo test -p windmill-ai anthropic`
Follow-up status: Bedrock native proxy handling has since moved into
`windmill-ai`, and the API-local `windmill-api/src/bedrock.rs` module has been
removed.
## Current Phase PR: Bedrock Native Proxy Migration
Goal: move the remaining native-provider API proxy execution out of
`windmill-api` and into `windmill-ai`, while leaving API-owned routing,
credential resolution, auditing, cache behavior, and Axum response conversion in
`windmill-api`.
Suggested PR title: `refactor(ai): move bedrock proxy handling to windmill-ai`.
Scope:
- Move Bedrock control-plane proxy calls (`foundation-models`,
`inference-profiles`) into `windmill-ai::providers::bedrock`.
- Move Bedrock chat proxy OpenAI request parsing, Converse request execution,
streaming SSE conversion, non-streaming OpenAI-shaped response conversion, and
auth selection into `windmill-ai::providers::bedrock`.
- Add an Axum-free `BedrockProxyResponse` shape in `windmill-ai`; the API route
converts it into an Axum body.
- Move the optional `aws-sdk-bedrock` dependency from `windmill-api` to
`windmill-ai`.
- Delete the API-local `windmill-api/src/bedrock.rs` module.
Out of scope:
- Do not unify `AIRequestConfig` and `ProviderWithResource`.
- Do not change Bedrock credential resolution, audit logging, request caching,
or non-Bedrock proxy behavior.
Validation:
- `cargo test -p windmill-ai bedrock --features bedrock`
- `cargo check -p windmill-ai -p windmill-api`
- `cargo check -p windmill-ai -p windmill-api --features bedrock`
## Known Follow-Ups
These are not blockers for the current migration PR because they either preserve
existing behavior or need a separate product decision, but they should stay
visible for later hardening work.
- **Google AI/Gemini native proxy custom headers**: the native Google AI proxy
path intentionally does not apply `AI_HTTP_HEADERS` or resource-level custom
headers today. Decide whether and how env/resource custom-header injection
should apply to Google AI once the proxy behavior is unified further.
- **Bedrock SSE tool-call indexing**: Bedrock streaming currently increments
the OpenAI tool-call index on every Bedrock `ContentBlockStop`, including text
content blocks. This behavior existed before the move from `windmill-api` to
`windmill-ai`, but a later cleanup should advance the index only when the
stopped block was a tool-use block.
- **Bedrock SSE keepalives**: Bedrock native SSE streams are still returned
directly without the API proxy keepalive injection used by other SSE paths.
This also preserves the pre-move behavior. A later cleanup can generalize the
keepalive wrapper so it works for both `reqwest::Error` streams and Bedrock's
SDK-backed `std::io::Error` streams.
## Step-by-Step Plan
Each step produces a compiling, working backend.
@@ -200,9 +256,10 @@ Move `AI_HTTP_HEADERS` lazy_static (currently duplicated in `windmill-api/src/ai
---
### Step 8: Add proxy support to QueryBuilder — API uses QueryBuilder for all providers
### Step 8: Add API proxy execution support to windmill-ai ✅
This is the key unification step. Add a new method to the `QueryBuilder` trait:
This is the key proxy unification step. HTTP-forwarding providers use
`QueryBuilder::build_proxy_request`:
```rust
/// Build a request from a raw OpenAI-format proxy request.
@@ -237,19 +294,21 @@ pub struct ProxyRequest {
**Provider implementations:**
- **OpenAI-compatible** (OpenAI, Mistral, DeepSeek, Groq, TogetherAI, CustomAI, OpenRouter): Minimal transformation — pass body through, build URL and auth headers.
- **Anthropic**: Handle standard vs Vertex AI. For Vertex: transform body (extract model, add anthropic_version). For standard: pass through with appropriate headers.
- **Google AI**: Convert OpenAI format → Gemini format (using existing `ai_google` functions). Replaces `windmill-api/src/google.rs`.
- **Bedrock**: Convert OpenAI format → Bedrock SDK calls. Replaces `windmill-api/src/bedrock.rs`.
- **Google AI**: Native execution mode converts OpenAI format → Gemini format and Gemini responses → OpenAI shape. Replaces `windmill-api/src/google.rs`.
- **Bedrock**: Native execution mode converts OpenAI format → Bedrock SDK calls and SDK responses → OpenAI shape. Replaces `windmill-api/src/bedrock.rs`.
**Refactor API proxy** (`windmill-api/src/ai.rs`):
1. Parse provider from headers, resolve credentials → `ProviderCredentials`
2. Create `QueryBuilder` via `create_query_builder`
3. Call `query_builder.build_proxy_request(&proxy_args)``ProxyRequest`
4. Send the request, return response with SSE keepalive injection
3. Dispatch by `ProxyExecutionMode`:
- HTTP-forwarding providers call `query_builder.build_proxy_request(&proxy_args)``ProxyRequest`
- Google AI and Bedrock call native handlers in `windmill-ai`
4. Convert the provider response to the API response body
**Remove** from windmill-api:
- `AIRequestConfig::prepare_request` — replaced by `QueryBuilder::build_proxy_request`
- `google.rs` — replaced by `GoogleAIQueryBuilder::build_proxy_request`
- `bedrock.rs` — replaced by `BedrockQueryBuilder::build_proxy_request`
- `google.rs` — replaced by `windmill_ai::providers::google_ai` native proxy handlers
- `bedrock.rs` — replaced by `windmill_ai::providers::bedrock` native proxy handlers
- `transform_anthropic_for_vertex` — moved to `AnthropicQueryBuilder`
- `supports_native_fim`, `transform_fim_to_chat_completions` — moved to windmill-ai
@@ -310,8 +369,8 @@ windmill-ai/src/
├── mod.rs # create_query_builder factory
├── anthropic.rs # build_request + build_proxy_request
├── openai.rs # build_request + build_proxy_request
├── google_ai.rs # build_request + build_proxy_request
├── bedrock.rs # build_request + build_proxy_request (feature: bedrock)
├── google_ai.rs # build_request + native proxy handlers
├── bedrock.rs # build_request + native proxy handlers (feature: bedrock)
├── other.rs # build_request + build_proxy_request
└── openrouter.rs # build_request + build_proxy_request
```