feat(ai): handle aws bedrock as provider (#7155)

* backend draft

* fix for tool and streaming

* do frontend side

* working

* working tools

* rm

* handle list endpoint

* handle for ai agents

* fix for models requiring inference id

* cleaning

* fix desc issue

* fix tool usage

* fix structured output

* cleaning

* fix for api

* rm

* fix input images

* cleaning

* chore: use aws sdk (#7156)

* feat(ai): Add AWS SDK dependencies for Bedrock integration

- Add aws-sdk-bedrockruntime v1.113.0
- Add aws-credential-types for bearer token authentication
- Update rustls to v0.23.35 for compatibility
- Dependencies added to windmill-common for AI features

* feat(ai): Add bearer token provider for Bedrock authentication

- Implement BearerTokenProvider using aws_credential_types
- Simple token-based auth using API keys from Windmill resources
- Add basic unit tests for provider creation
- Export bedrock_auth module in lib.rs

* feat(ai): Add Bedrock client wrapper with region extraction

- Implement BedrockClient wrapper around AWS SDK client
- Bearer token authentication integration
- Extract AWS region from Bedrock base URL automatically
- Comprehensive unit tests for region extraction
- Make aws-config non-optional dependency for AI features
- Update feature flags to reflect new dependency structure

* cargo

* feat(ai): Implement non-streaming Bedrock via AWS SDK

Use official AWS SDK instead of manual HTTP requests for better type safety and maintainability. Implements the Bedrock converse() API for non-streaming requests with proper bearer token authentication and message format conversion between OpenAI and Bedrock formats.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>

* refactor(ai): Eliminate Simple* conversion types for Bedrock SDK

- Move AI types to windmill-common/src/ai_types.rs for shared access
- Update bedrock_converters to work directly with OpenAI types
- Remove ~200 lines of conversion boilerplate from ai_executor.rs and bedrock.rs
- Remove unused imports to clean compilation warnings
- Benefits: 50% fewer conversion steps, no information loss, easier maintenance

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>

* feat(ai): Add streaming support for AWS Bedrock SDK

- Implement converse_stream() for Bedrock streaming responses
- Use EventReceiver.recv() to process stream events
- Extract text deltas using bedrock_stream_event_to_text()
- Send TokenDelta events to StreamEventProcessor for real-time updates
- Refactor request building to eliminate duplication between streaming and non-streaming
- Clean, minimal implementation following AWS SDK patterns

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>

* revert flake change

* fix

* feat(ai): Add tool calls and image support for Bedrock streaming

**Phase 1: Streaming Tool Call Support**
- Add stream event processing functions in bedrock_converters.rs:
  - bedrock_stream_event_to_tool_start() - Extract tool use start from ContentBlockStart
  - bedrock_stream_event_to_tool_delta() - Extract tool input deltas from ContentBlockDelta
  - bedrock_stream_event_is_block_stop() - Detect ContentBlockStop events
  - streaming_tool_calls_to_openai() - Convert accumulated tool calls to OpenAI format
- Update ai_executor.rs streaming loop with tool call accumulator (HashMap)
- Track current tool use ID during streaming
- Send ToolCallArguments events to StreamEventProcessor
- Return accumulated tool calls instead of empty vector

**Phase 2: Image Input Support**
- Add parse_image_data_url() to extract format and base64 data from data URLs
- Add content_part_to_block() to convert ContentPart to Bedrock ContentBlock
- Refactor convert_message() to handle multi-part content with images
- Support ImageUrl conversion to Bedrock ImageBlock with proper format (png/jpeg/gif/webp)
- Import AWS SDK image types: ImageBlock, ImageSource, ImageFormat
- Keep content_to_text() helper for system message text extraction

**Benefits**:
-  Tool calling now works in both streaming and non-streaming modes
-  Images are properly converted instead of being silently dropped
-  Structured output works in streaming (uses tool calling)
-  Full feature parity with manual HTTP implementation

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>

* cleaning

* fix(ai): Add S3 image support and structured output for Bedrock

**Fixes:**
1. **S3 Image Support**: Call prepare_messages_for_api() before Bedrock SDK path to convert S3Objects to ImageUrls
   - Downloads images from S3 and encodes as base64 data URLs
   - Ensures images are properly handled in both streaming and non-streaming modes

2. **Structured Output**: Add ToolChoice::Any when structured output tool is present
   - Forces Bedrock to call the structured_output tool
   - Ensures JSON schema compliance for structured output
   - Works in both streaming and non-streaming modes

**Changes:**
- ai_executor.rs: Call prepare_messages_for_api() for Bedrock SDK path
- ai_executor.rs: Set tool_choice to Any when structured_output_tool_name is present
- aws_bedrock.rs: Remove unused ToolChoice imports (used via full path in worker)

**Testing:**
-  S3 images are now downloaded and converted before API call
-  Structured output now forces tool usage with ToolChoice::Any
-  Both work in streaming and non-streaming modes

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>

* cleaning

* cleaning

* cleaning

* better error

* cleaning

* cleaning

* rm

* rename

* apply region

---------

Co-authored-by: Claude <noreply@anthropic.com>

* fix default

* no panic

* no print

* use utils file

* cleaning

---------

Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
centdix
2025-11-17 19:57:59 +01:00
committed by GitHub
parent da4f57ae59
commit 79ac6312e8
21 changed files with 2109 additions and 391 deletions
+13
View File
@@ -10,3 +10,16 @@
1. Update database schema with migration if necessary
2. Update backend/windmill-api/openapi.yaml after modifying API endpoints
## Querying the Database
To query the database directly, use psql with the following connection string:
```bash
psql postgres://postgres:changeme@localhost:5432/windmill
```
This can be helpful for:
- Inspecting database state during development
- Testing queries before implementing them in Rust
- Debugging data-related issues
+68 -22
View File
@@ -823,13 +823,14 @@ dependencies = [
[[package]]
name = "aws-runtime"
version = "1.5.10"
version = "1.5.14"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c034a1bc1d70e16e7f4e4caf7e9f7693e4c9c24cd91cf17c2a0b21abaebc7c8b"
checksum = "8fe0fd441565b0b318c76e7206c8d1d0b0166b3e986cf30e890b61feb6192045"
dependencies = [
"aws-credential-types",
"aws-sigv4",
"aws-smithy-async",
"aws-smithy-eventstream",
"aws-smithy-http",
"aws-smithy-runtime",
"aws-smithy-runtime-api",
@@ -845,6 +846,31 @@ dependencies = [
"uuid",
]
[[package]]
name = "aws-sdk-bedrockruntime"
version = "1.113.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d5d2b8f081b9e8ff455b8dd7387b6b02263c3dac73172d188d2b523ff1e775e9"
dependencies = [
"aws-credential-types",
"aws-runtime",
"aws-sigv4",
"aws-smithy-async",
"aws-smithy-eventstream",
"aws-smithy-http",
"aws-smithy-json",
"aws-smithy-runtime",
"aws-smithy-runtime-api",
"aws-smithy-types",
"aws-types",
"bytes",
"fastrand",
"http 0.2.12",
"hyper 0.14.32",
"regex-lite",
"tracing",
]
[[package]]
name = "aws-sdk-config"
version = "1.68.0"
@@ -964,6 +990,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c35452ec3f001e1f2f6db107b6373f1f48f05ec63ba2c5c9fa91f07dad32af11"
dependencies = [
"aws-credential-types",
"aws-smithy-eventstream",
"aws-smithy-http",
"aws-smithy-runtime-api",
"aws-smithy-types",
@@ -990,12 +1017,24 @@ dependencies = [
"tokio",
]
[[package]]
name = "aws-smithy-eventstream"
version = "0.60.13"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e29a304f8319781a39808847efb39561351b1bb76e933da7aa90232673638658"
dependencies = [
"aws-smithy-types",
"bytes",
"crc32fast",
]
[[package]]
name = "aws-smithy-http"
version = "0.62.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "445d5d720c99eed0b4aa674ed00d835d9b1427dd73e04adaf2f94c6b2d6f9fca"
dependencies = [
"aws-smithy-eventstream",
"aws-smithy-runtime-api",
"aws-smithy-types",
"bytes",
@@ -1013,9 +1052,9 @@ dependencies = [
[[package]]
name = "aws-smithy-http-client"
version = "1.0.6"
version = "1.1.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f108f1ca850f3feef3009bdcc977be201bca9a91058864d9de0684e64514bee0"
checksum = "623254723e8dfd535f566ee7b2381645f8981da086b5c4aa26c0c41582bb1d2c"
dependencies = [
"aws-smithy-async",
"aws-smithy-runtime-api",
@@ -1032,10 +1071,11 @@ dependencies = [
"hyper-util",
"pin-project-lite",
"rustls 0.21.12",
"rustls 0.23.29",
"rustls 0.23.35",
"rustls-native-certs 0.8.2",
"rustls-pki-types",
"tokio",
"tokio-rustls 0.26.4",
"tower 0.5.2",
"tracing",
]
@@ -1070,9 +1110,9 @@ dependencies = [
[[package]]
name = "aws-smithy-runtime"
version = "1.8.6"
version = "1.9.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9e107ce0783019dbff59b3a244aa0c114e4a8c9d93498af9162608cd5474e796"
checksum = "0bbe9d018d646b96c7be063dd07987849862b0e6d07c778aad7d93d1be6c1ef0"
dependencies = [
"aws-smithy-async",
"aws-smithy-http",
@@ -4266,7 +4306,7 @@ dependencies = [
"deno_core",
"deno_error",
"deno_native_certs",
"rustls 0.23.29",
"rustls 0.23.35",
"rustls-pemfile 2.2.0",
"rustls-tokio-stream",
"rustls-webpki 0.102.8",
@@ -6686,7 +6726,7 @@ dependencies = [
"hyper 1.8.1",
"hyper-util",
"log",
"rustls 0.23.29",
"rustls 0.23.35",
"rustls-native-certs 0.8.2",
"rustls-pki-types",
"tokio",
@@ -7428,7 +7468,7 @@ dependencies = [
"k8s-openapi",
"kube-core",
"pem 3.0.5",
"rustls 0.23.29",
"rustls 0.23.35",
"secrecy",
"serde",
"serde_json",
@@ -7913,7 +7953,7 @@ dependencies = [
"base64 0.22.1",
"gethostname",
"mail-builder",
"rustls 0.23.29",
"rustls 0.23.35",
"rustls-pki-types",
"smtp-proto",
"tokio",
@@ -10262,7 +10302,7 @@ dependencies = [
"quinn-proto",
"quinn-udp",
"rustc-hash 2.1.1",
"rustls 0.23.29",
"rustls 0.23.35",
"socket2 0.6.1",
"thiserror 2.0.17",
"tokio",
@@ -10282,7 +10322,7 @@ dependencies = [
"rand 0.9.0",
"ring 0.17.14",
"rustc-hash 2.1.1",
"rustls 0.23.29",
"rustls 0.23.35",
"rustls-pki-types",
"slab",
"thiserror 2.0.17",
@@ -10724,7 +10764,7 @@ dependencies = [
"percent-encoding",
"pin-project-lite",
"quinn",
"rustls 0.23.29",
"rustls 0.23.35",
"rustls-native-certs 0.8.2",
"rustls-pki-types",
"serde",
@@ -11146,9 +11186,9 @@ dependencies = [
[[package]]
name = "rustls"
version = "0.23.29"
version = "0.23.35"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2491382039b29b9b11ff08b76ff6c97cf287671dbb74f0be44bda389fffe9bd1"
checksum = "533f54bc6a7d4f647e46ad909549eda97bf5afc1585190ef692b4286b198bd8f"
dependencies = [
"aws-lc-rs",
"log",
@@ -11232,7 +11272,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "22557157d7395bc30727745b365d923f1ecc230c4c80b176545f3f4f08c46e33"
dependencies = [
"futures",
"rustls 0.23.29",
"rustls 0.23.35",
"socket2 0.5.10",
"tokio",
]
@@ -12297,7 +12337,7 @@ dependencies = [
"memchr",
"once_cell",
"percent-encoding",
"rustls 0.23.29",
"rustls 0.23.35",
"serde",
"serde_json",
"sha2 0.10.9",
@@ -13771,7 +13811,7 @@ version = "0.26.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61"
dependencies = [
"rustls 0.23.29",
"rustls 0.23.35",
"tokio",
]
@@ -14603,7 +14643,7 @@ dependencies = [
"log",
"native-tls",
"once_cell",
"rustls 0.23.29",
"rustls 0.23.35",
"rustls-pki-types",
"serde",
"serde_json",
@@ -15177,7 +15217,7 @@ dependencies = [
"quote",
"rand 0.9.0",
"reqwest 0.12.24",
"rustls 0.23.29",
"rustls 0.23.35",
"serde",
"serde_json",
"serde_yml",
@@ -15283,7 +15323,7 @@ dependencies = [
"rumqttc",
"rust-embed",
"rust_decimal",
"rustls 0.23.29",
"rustls 0.23.35",
"samael",
"serde",
"serde_json",
@@ -15386,7 +15426,9 @@ dependencies = [
"async-stream",
"async-trait",
"aws-config",
"aws-credential-types",
"aws-sdk-sts",
"aws-smithy-types",
"aws-smithy-types-convert",
"axum",
"backon",
@@ -15768,6 +15810,10 @@ dependencies = [
"async-recursion",
"async-stream",
"async-trait",
"aws-config",
"aws-credential-types",
"aws-sdk-bedrockruntime",
"aws-smithy-types",
"backon",
"base64 0.22.1",
"bit-vec 0.6.3",
+4 -1
View File
@@ -384,11 +384,14 @@ datafusion = "47.0.0"
object_store = { git = "https://github.com/apache/arrow-rs-object-store", rev = "36752c975d4f29e20b57c91f81a10872dcd48ae7", features = ["aws", "azure", "gcp"] }
openidconnect = { version = "4.0.0-rc.1" }
aws-config = "^1"
aws-sdk-bedrockruntime = "=1.113.0"
aws-credential-types = "^1"
aws-smithy-types = "^1"
aws-sdk-sqs = "=1.77.0"
aws-sdk-sts = "=1.79.0"
aws-sdk-sso = "=1.77.0"
aws-sdk-ssooidc = "=1.78.0"
rustls = "=0.23.29"
rustls = "=0.23.35"
async-once-cell = "0.5.4"
+29 -30
View File
@@ -572,12 +572,12 @@ paths:
use_case:
type: string
responses:
'200':
"200":
description: Onboarding data submitted successfully
content:
application/json:
schema:
type: string
type: string
/w/{workspace}/users/delete/{username}:
delete:
@@ -15294,8 +15294,7 @@ components:
CreatedAfterQueue:
name: created_after_queue
description:
filter on jobs created after X for jobs in the queue only
description: filter on jobs created after X for jobs in the queue only
in: query
schema:
type: string
@@ -15303,8 +15302,7 @@ components:
CreatedBeforeQueue:
name: created_before_queue
description:
filter on jobs created before X for jobs in the queue only
description: filter on jobs created before X for jobs in the queue only
in: query
schema:
type: string
@@ -15599,6 +15597,7 @@ components:
groq,
openrouter,
togetherai,
aws_bedrock,
customai,
]
@@ -16750,30 +16749,30 @@ components:
ScriptLang:
type: string
enum: [
python3,
deno,
go,
bash,
powershell,
postgresql,
mysql,
bigquery,
snowflake,
mssql,
oracledb,
graphql,
nativets,
bun,
php,
rust,
ansible,
csharp,
nu,
java,
ruby,
duckdb,
# for related places search: ADD_NEW_LANG
]
python3,
deno,
go,
bash,
powershell,
postgresql,
mysql,
bigquery,
snowflake,
mssql,
oracledb,
graphql,
nativets,
bun,
php,
rust,
ansible,
csharp,
nu,
java,
ruby,
duckdb,
# for related places search: ADD_NEW_LANG
]
Preview:
type: object
+89 -12
View File
@@ -1,3 +1,4 @@
use crate::bedrock;
use crate::db::{ApiAuthed, DB};
use axum::{body::Bytes, extract::Path, response::IntoResponse, routing::post, Extension, Router};
@@ -6,12 +7,12 @@ use quick_cache::sync::Cache;
use reqwest::{Client, RequestBuilder};
use serde::{Deserialize, Serialize};
use serde_json::value::RawValue;
use windmill_common::variables::get_variable_or_self;
use std::collections::HashMap;
use windmill_audit::{audit_oss::audit_log, ActionKind};
use windmill_common::ai_providers::{AIProvider, ProviderConfig, ProviderModel, AZURE_API_VERSION};
use windmill_common::error::{to_anyhow, Error, Result};
use windmill_common::utils::configure_client;
use windmill_common::variables::get_variable_or_self;
lazy_static::lazy_static! {
static ref HTTP_CLIENT: Client = configure_client(reqwest::ClientBuilder::new()
@@ -66,6 +67,7 @@ struct AIStandardResource {
#[serde(alias = "apiKey")]
api_key: Option<String>,
organization_id: Option<String>,
region: Option<String>,
}
#[derive(Deserialize, Debug)]
@@ -98,7 +100,9 @@ impl AIRequestConfig {
) -> Result<Self> {
let (api_key, access_token, organization_id, base_url, user) = match resource {
AIResource::Standard(resource) => {
let base_url = provider.get_base_url(resource.base_url, db).await?;
let base_url = provider
.get_base_url(resource.base_url, resource.region, db)
.await?;
let api_key = if let Some(api_key) = resource.api_key {
Some(get_variable_or_self(api_key, db, w_id).await?)
} else {
@@ -119,7 +123,7 @@ impl AIRequestConfig {
None
};
let token = Self::get_token_using_oauth(resource, db, w_id).await?;
let base_url = provider.get_base_url(None, db).await?;
let base_url = provider.get_base_url(None, None, db).await?;
(None, Some(token), None, base_url, user)
}
@@ -180,15 +184,35 @@ impl AIRequestConfig {
let is_azure = provider.is_azure_openai(base_url);
let is_anthropic = matches!(provider, AIProvider::Anthropic);
let is_anthropic_sdk = headers.get("X-Anthropic-SDK").is_some();
let is_bedrock = matches!(provider, AIProvider::AWSBedrock);
let url = if is_azure && method != Method::GET {
// Handle AWS Bedrock transformation
let (url, body) = if is_bedrock && method != Method::GET {
let (model, transformed_body, is_streaming) =
bedrock::transform_openai_to_bedrock(&body)?;
let endpoint = if is_streaming {
"converse-stream"
} else {
"converse"
};
let bedrock_url = format!("{}/model/{}/{}", base_url, model, endpoint);
(bedrock_url, transformed_body)
} else if is_bedrock && (path == "foundation-models" || path == "inference-profiles") {
// AWS Bedrock foundation-models and inference-profiles endpoints use different base URL (without -runtime)
let bedrock_base_url = base_url.replace("bedrock-runtime.", "bedrock.");
let bedrock_url = format!("{}/{}", bedrock_base_url, path);
(bedrock_url, body)
} else if is_azure && method != Method::GET {
let model = AIProvider::extract_model_from_body(&body)?;
AIProvider::build_azure_openai_url(base_url, &model, path)
let azure_url = AIProvider::build_azure_openai_url(base_url, &model, path);
(azure_url, body)
} else if is_anthropic_sdk {
let truncated_base_url = base_url.trim_end_matches("/v1");
format!("{}/{}", truncated_base_url, path)
let anthropic_url = format!("{}/{}", truncated_base_url, path);
(anthropic_url, body)
} else {
format!("{}/{}", base_url, path)
let default_url = format!("{}/{}", base_url, path);
(default_url, body)
};
tracing::debug!("AI request URL: {}", url);
@@ -316,7 +340,7 @@ async fn global_proxy(
return Err(Error::BadRequest("API key is required".to_string()));
};
let base_url = provider.get_base_url(None, &db).await?;
let base_url = provider.get_base_url(None, None, &db).await?;
let url = format!("{}/{}", base_url, ai_path);
@@ -446,6 +470,21 @@ async fn proxy(
}
};
// Extract model and streaming flag for Bedrock transformation (only for POST requests)
let (model, is_streaming) =
if matches!(provider, AIProvider::AWSBedrock) && method == Method::POST {
#[derive(Deserialize, Debug)]
struct BedrockRequest {
model: String,
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)
};
let request = request_config.prepare_request(&provider, &ai_path, method, headers, body)?;
let response = request.send().await.map_err(to_anyhow)?;
@@ -469,8 +508,46 @@ async fn proxy(
return Err(Error::AIError(err_msg));
}
let status_code = response.status();
let headers = response.headers().clone();
let stream = response.bytes_stream();
Ok((status_code, headers, axum::body::Body::from_stream(stream)))
// Transform Bedrock responses back to OpenAI format
if matches!(provider, AIProvider::AWSBedrock) {
let Some(model) = model else {
return Err(Error::BadRequest("Model is required".to_string()));
};
if is_streaming {
// Transform streaming response
use http::StatusCode;
let mut response_headers = 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());
let stream = response.bytes_stream();
let transformed_stream = bedrock::transform_bedrock_stream_to_openai(stream, model);
Ok((
StatusCode::OK,
response_headers,
axum::body::Body::from_stream(transformed_stream),
))
} else {
// Transform non-streaming response
let transformed_body = bedrock::transform_bedrock_to_openai(response, model).await?;
let mut response_headers = HeaderMap::new();
response_headers.insert("content-type", "application/json".parse().unwrap());
Ok((
http::StatusCode::OK,
response_headers,
axum::body::Body::from(transformed_body),
))
}
} else {
// Pass through for other providers
let status_code = response.status();
let headers = response.headers().clone();
let stream = response.bytes_stream();
Ok((status_code, headers, axum::body::Body::from_stream(stream)))
}
}
+602
View File
@@ -0,0 +1,602 @@
use axum::body::Bytes;
use bytes;
use futures;
use uuid;
use windmill_common::error::{Error, Result};
/// Transform OpenAI format request to AWS Bedrock Converse format
/// Returns: (model_id, transformed_body, is_streaming)
pub fn transform_openai_to_bedrock(body: &[u8]) -> Result<(String, Bytes, bool)> {
use serde_json::Value;
// Parse the OpenAI request
let openai_req: Value = serde_json::from_slice(body)
.map_err(|e| Error::internal_err(format!("Failed to parse OpenAI request: {}", e)))?;
// Extract model and streaming flag
let model = openai_req["model"]
.as_str()
.ok_or_else(|| Error::BadRequest("Missing 'model' field in request".to_string()))?
.to_string();
let is_streaming = openai_req["stream"].as_bool().unwrap_or(false);
// Build Bedrock request
let mut bedrock_req = serde_json::json!({});
// Transform messages
if let Some(messages) = openai_req["messages"].as_array() {
let mut system_messages = Vec::new();
let mut conversation_messages = Vec::new();
for msg in messages {
let role = msg["role"].as_str().unwrap_or("");
match role {
"system" => {
// Extract system messages to separate array
if let Some(content) = msg["content"].as_str() {
system_messages.push(serde_json::json!({"text": content}));
}
}
"user" | "assistant" => {
// Normalize content to array format
let mut content = if let Some(text) = msg["content"].as_str() {
// Simple string → array of content blocks
vec![serde_json::json!({"text": text})]
} else if let Some(content_array) = msg["content"].as_array() {
// Already an array - transform each item
content_array
.iter()
.filter_map(|item| {
if let Some(text) = item["text"].as_str() {
Some(serde_json::json!({"text": text}))
} else if item["type"].as_str() == Some("text") {
Some(serde_json::json!({"text": item["text"]}))
} else if item["type"].as_str() == Some("image_url") {
// Transform image_url format if needed
// For now, pass through - may need more sophisticated handling
Some(item.clone())
} else {
None
}
})
.collect()
} else {
vec![]
};
// Handle tool_calls for assistant messages (OpenAI → Bedrock toolUse)
if role == "assistant" {
if let Some(tool_calls) = msg["tool_calls"].as_array() {
for tool_call in tool_calls {
if tool_call["type"].as_str() == Some("function") {
let tool_use_id = tool_call["id"].as_str().unwrap_or("");
let function_name =
tool_call["function"]["name"].as_str().unwrap_or("");
let arguments_str =
tool_call["function"]["arguments"].as_str().unwrap_or("{}");
// Parse arguments JSON string to object
let input = serde_json::from_str::<Value>(arguments_str)
.map_err(|e| {
Error::internal_err(format!(
"Failed to parse tool call arguments: {}",
e
))
})?;
content.push(serde_json::json!({
"toolUse": {
"toolUseId": tool_use_id,
"name": function_name,
"input": input
}
}));
}
}
}
}
// Only add message if it has content
if !content.is_empty() {
conversation_messages.push(serde_json::json!({
"role": role,
"content": content
}));
}
}
"tool" => {
// Transform tool response to Bedrock format
let tool_call_id = msg["tool_call_id"].as_str().unwrap_or("");
let content = msg["content"].as_str().unwrap_or("");
// Try to parse content as JSON
// Bedrock requires json field to be an object, not a primitive or array
let tool_result_content =
if let Ok(json_content) = serde_json::from_str::<Value>(content) {
if json_content.is_object() {
vec![serde_json::json!({"json": json_content})]
} else {
// Wrap primitives and arrays in an object
vec![serde_json::json!({"json": {"result": json_content}})]
}
} else {
vec![serde_json::json!({"text": content})]
};
conversation_messages.push(serde_json::json!({
"role": "user",
"content": [{
"toolResult": {
"toolUseId": tool_call_id,
"content": tool_result_content
}
}]
}));
}
_ => {}
}
}
if !system_messages.is_empty() {
bedrock_req["system"] = Value::Array(system_messages);
}
bedrock_req["messages"] = Value::Array(conversation_messages);
}
// Transform inference parameters
let mut inference_config = serde_json::json!({});
if let Some(max_tokens) = openai_req["max_tokens"].as_i64() {
inference_config["maxTokens"] = Value::Number(max_tokens.into());
}
if let Some(temperature) = openai_req["temperature"].as_f64() {
inference_config["temperature"] = serde_json::json!(temperature);
}
if let Some(top_p) = openai_req["top_p"].as_f64() {
inference_config["topP"] = serde_json::json!(top_p);
}
if let Some(stop) = openai_req["stop"].as_array() {
let stop_sequences: Vec<String> = stop
.iter()
.filter_map(|s| s.as_str().map(|s| s.to_string()))
.collect();
if !stop_sequences.is_empty() {
inference_config["stopSequences"] =
Value::Array(stop_sequences.into_iter().map(Value::String).collect());
}
}
if !inference_config.as_object().unwrap().is_empty() {
bedrock_req["inferenceConfig"] = inference_config;
}
// Transform tools if present
if let Some(tools) = openai_req["tools"].as_array() {
let mut bedrock_tools = Vec::new();
for tool in tools {
if tool["type"].as_str() == Some("function") {
if let Some(function) = tool["function"].as_object() {
bedrock_tools.push(serde_json::json!({
"toolSpec": {
"name": function.get("name"),
"description": function.get("description")
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
.unwrap_or("Tool function"),
"inputSchema": {
"json": function.get("parameters")
}
}
}));
}
}
}
if !bedrock_tools.is_empty() {
let mut tool_config = serde_json::json!({
"tools": bedrock_tools
});
// Transform tool_choice
if let Some(tool_choice) = openai_req.get("tool_choice") {
if tool_choice == "auto" {
tool_config["toolChoice"] = serde_json::json!({"auto": {}});
} else if tool_choice == "required" {
tool_config["toolChoice"] = serde_json::json!({"any": {}});
} else if let Some(obj) = tool_choice.as_object() {
if obj.get("type").and_then(|v| v.as_str()) == Some("function") {
if let Some(function) = obj.get("function").and_then(|v| v.as_object()) {
if let Some(name) = function.get("name").and_then(|v| v.as_str()) {
tool_config["toolChoice"] = serde_json::json!({
"tool": {"name": name}
});
}
}
}
}
}
bedrock_req["toolConfig"] = tool_config;
}
}
let transformed_body = serde_json::to_vec(&bedrock_req)
.map_err(|e| Error::internal_err(format!("Failed to serialize Bedrock request: {}", e)))?
.into();
Ok((model, transformed_body, is_streaming))
}
/// Transform AWS Bedrock Converse response to OpenAI format
pub async fn transform_bedrock_to_openai(
response: reqwest::Response,
model: String,
) -> Result<Bytes> {
use serde_json::Value;
let bedrock_resp: Value = response
.json()
.await
.map_err(|e| Error::internal_err(format!("Failed to parse Bedrock response: {}", e)))?;
// Generate unique ID and timestamp
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 and map to finish_reason
let stop_reason = bedrock_resp["stopReason"].as_str().unwrap_or("end_turn");
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 message_content = &bedrock_resp["output"]["message"]["content"];
let mut text_content = String::new();
let mut tool_calls = Vec::new();
if let Some(content_array) = message_content.as_array() {
for (_index, block) in content_array.iter().enumerate() {
if let Some(text) = block["text"].as_str() {
text_content.push_str(text);
} else if let Some(tool_use) = block.get("toolUse") {
// Transform tool use to OpenAI tool_calls format
let tool_call_id = tool_use["toolUseId"].as_str().unwrap_or("");
let name = tool_use["name"].as_str().unwrap_or("");
let input = &tool_use["input"];
tool_calls.push(serde_json::json!({
"id": tool_call_id,
"type": "function",
"function": {
"name": name,
"arguments": serde_json::to_string(input).unwrap_or_default()
}
}));
}
}
}
// Build the message
let message = if !tool_calls.is_empty() {
serde_json::json!({
"role": "assistant",
"content": if text_content.is_empty() { Value::Null } else { 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) = bedrock_resp.get("usage") {
serde_json::json!({
"prompt_tokens": usage_data["inputTokens"].as_i64().unwrap_or(0),
"completion_tokens": usage_data["outputTokens"].as_i64().unwrap_or(0),
"total_tokens": usage_data["totalTokens"].as_i64().unwrap_or(0)
})
} 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)))?
.into();
Ok(response_body)
}
/// Transform AWS Bedrock streaming response to OpenAI SSE format
/// Bedrock uses AWS event stream binary format, not SSE
pub fn transform_bedrock_stream_to_openai(
stream: impl futures::Stream<Item = std::result::Result<bytes::Bytes, reqwest::Error>>
+ Send
+ 'static,
model: String,
) -> impl futures::Stream<Item = std::result::Result<bytes::Bytes, std::io::Error>> + Send {
use futures::stream::StreamExt;
use serde_json::Value;
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 and binary buffer
struct StreamState {
id: String,
model: String,
created: u64,
tool_calls: HashMap<usize, (String, String, String)>, // index -> (id, name, args)
buffer: Vec<u8>, // Binary buffer for AWS event stream
}
let state = std::sync::Arc::new(tokio::sync::Mutex::new(StreamState {
id: id.clone(),
model: model.clone(),
created,
tool_calls: HashMap::new(),
buffer: Vec::new(),
}));
stream
.then(move |chunk_result| {
let state = state.clone();
async move {
match chunk_result {
Ok(chunk) => {
let mut state = state.lock().await;
state.buffer.extend_from_slice(&chunk);
let mut events = Vec::new();
// Parse AWS event stream messages from buffer
loop {
// Need at least 12 bytes for prelude (8) + prelude CRC (4)
if state.buffer.len() < 12 {
break;
}
// Read prelude: total_length (4 bytes) + headers_length (4 bytes)
let total_length = u32::from_be_bytes([
state.buffer[0],
state.buffer[1],
state.buffer[2],
state.buffer[3],
]) as usize;
// Check if we have the complete message
if state.buffer.len() < total_length {
break;
}
let headers_length = u32::from_be_bytes([
state.buffer[4],
state.buffer[5],
state.buffer[6],
state.buffer[7],
]) as usize;
// Skip prelude CRC (4 bytes after prelude)
let headers_start = 12;
let payload_start = headers_start + headers_length;
let payload_end = total_length - 4; // Exclude message CRC
// Parse headers to extract event type
let mut event_type = None;
let mut pos = headers_start;
while pos < payload_start {
if pos + 1 > state.buffer.len() {
break;
}
let name_len = state.buffer[pos] as usize;
pos += 1;
if pos + name_len > state.buffer.len() {
break;
}
let name = String::from_utf8_lossy(&state.buffer[pos..pos + name_len]).to_string();
pos += name_len;
if pos + 3 > state.buffer.len() {
break;
}
let value_type = state.buffer[pos];
pos += 1;
let value_len = u16::from_be_bytes([state.buffer[pos], state.buffer[pos + 1]]) as usize;
pos += 2;
if pos + value_len > state.buffer.len() {
break;
}
if value_type == 7 && name == ":event-type" {
event_type = Some(String::from_utf8_lossy(&state.buffer[pos..pos + value_len]).to_string());
}
pos += value_len;
}
// Extract JSON payload (copy to avoid borrow issues)
let payload = state.buffer[payload_start..payload_end].to_vec();
// Remove processed message from buffer
state.buffer.drain(0..total_length);
// Process the event
if let Some(evt_type) = event_type {
if let Ok(payload_str) = std::str::from_utf8(&payload) {
if let Ok(parsed_data) = serde_json::from_str::<Value>(payload_str) {
// Transform based on event type
match evt_type.as_str() {
"messageStart" => {
// No output for messageStart
}
"contentBlockStart" => {
let index = parsed_data["contentBlockIndex"].as_u64().unwrap_or(0) as usize;
if let Some(tool_use) = parsed_data["start"].get("toolUse") {
let tool_id = tool_use["toolUseId"].as_str().unwrap_or("").to_string();
let name = tool_use["name"].as_str().unwrap_or("").to_string();
state.tool_calls.insert(index, (tool_id.clone(), 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_id,
"type": "function",
"function": {
"name": name,
"arguments": ""
}
}]
},
"finish_reason": Value::Null
}]
});
events.push(Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk))));
}
}
"contentBlockDelta" => {
let index = parsed_data["contentBlockIndex"].as_u64().unwrap_or(0) as usize;
if let Some(text) = parsed_data["delta"]["text"].as_str() {
// Text content delta
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": Value::Null
}]
});
events.push(Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk))));
} else if let Some(tool_use_input) = parsed_data["delta"]["toolUse"]["input"].as_str() {
// Tool use arguments delta
if let Some((_tool_id, _name, ref mut args)) = state.tool_calls.get_mut(&index) {
args.push_str(tool_use_input);
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": tool_use_input
}
}]
},
"finish_reason": Value::Null
}]
});
events.push(Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk))));
}
}
}
"contentBlockStop" => {
// No output needed
}
"messageStop" => {
let stop_reason = parsed_data["stopReason"].as_str().unwrap_or("end_turn");
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
}]
});
events.push(Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk))));
}
"metadata" => {
// Could include usage info here if needed
}
_ => {}
}
}
}
}
} // end loop
events
}
Err(e) => {
vec![Err(std::io::Error::new(
std::io::ErrorKind::Other,
e.to_string(),
))]
}
}
}
})
.flat_map(|events| futures::stream::iter(events))
.chain(futures::stream::iter(vec![
// Send [DONE] at the end
Ok(bytes::Bytes::from("data: [DONE]\n\n"))
]))
}
+1
View File
@@ -75,6 +75,7 @@ pub mod agent_workers_ee;
#[cfg(feature = "agent_worker_server")]
mod agent_workers_oss;
mod ai;
mod bedrock;
mod apps;
pub mod args;
mod assets;
+2
View File
@@ -64,6 +64,8 @@ object_store = { workspace = true, optional = true }
prometheus = { workspace = true, optional = true }
aws-config = { workspace = true, optional = true }
aws-sdk-sts = { workspace = true, optional = true }
aws-credential-types.workspace = true
aws-smithy-types.workspace = true
base64.workspace = true
bitflags.workspace = true
+12 -1
View File
@@ -27,11 +27,18 @@ pub enum AIProvider {
OpenRouter,
TogetherAI,
CustomAI,
#[serde(rename = "aws_bedrock")]
AWSBedrock,
}
impl AIProvider {
/// Get the base URL for the AI provider
pub async fn get_base_url(&self, resource_base_url: Option<String>, db: &DB) -> Result<String> {
pub async fn get_base_url(
&self,
resource_base_url: Option<String>,
region: Option<String>,
db: &DB,
) -> Result<String> {
match self {
AIProvider::OpenAI => {
// Check for Azure base path override
@@ -74,6 +81,10 @@ impl AIProvider {
)))
}
}
AIProvider::AWSBedrock => Ok(format!(
"https://bedrock-runtime.{}.amazonaws.com",
region.unwrap_or_else(|| "us-east-1".to_string())
)),
}
}
+4
View File
@@ -58,6 +58,10 @@ windmill-parser-graphql.workspace = true
windmill-parser-php = { workspace = true, optional = true }
windmill-git-sync.workspace = true
rmcp = { version = "0.8.1", features = ["client", "transport-streamable-http-client", "transport-streamable-http-client-reqwest"] }
aws-sdk-bedrockruntime.workspace = true
aws-config.workspace = true
aws-credential-types.workspace = true
aws-smithy-types.workspace = true
flume.workspace = true
sqlx.workspace = true
uuid.workspace = true
@@ -4,6 +4,8 @@ use ulid;
use windmill_common::{client::AuthedClient, error::Error, s3_helpers::S3Object};
use windmill_queue::MiniPulledJob;
use crate::ai::types::*;
/// Upload image to S3 and return S3Object
pub async fn upload_image_to_s3(
base64_image: &str,
@@ -66,3 +68,53 @@ pub async fn download_and_encode_s3_image(
Ok((mime_type.to_string(), base64_data))
}
/// Prepare messages for API by converting S3Objects to base64 ImageUrls
pub async fn prepare_messages_for_api(
messages: &[OpenAIMessage],
client: &AuthedClient,
workspace_id: &str,
) -> Result<Vec<OpenAIMessage>, Error> {
let mut prepared_messages = Vec::new();
for message in messages {
let mut prepared_message = message.clone();
if let Some(content) = &message.content {
match content {
OpenAIContent::Text(text) => {
prepared_message.content = Some(OpenAIContent::Text(text.clone()));
}
OpenAIContent::Parts(parts) => {
let mut prepared_content = Vec::new();
for part in parts {
match part {
ContentPart::S3Object { s3_object } => {
// Convert S3Object to base64 image URL
let (mime_type, image_bytes) =
download_and_encode_s3_image(s3_object, client, workspace_id)
.await?;
prepared_content.push(ContentPart::ImageUrl {
image_url: ImageUrlData {
url: format!("data:{};base64,{}", mime_type, image_bytes),
},
});
}
other => {
// Keep Text and ImageUrl as-is
prepared_content.push(other.clone());
}
}
}
prepared_message.content = Some(OpenAIContent::Parts(prepared_content));
}
}
}
prepared_messages.push(prepared_message);
}
Ok(prepared_messages)
}
@@ -0,0 +1,804 @@
use crate::ai::{
image_handler::prepare_messages_for_api,
providers::openai::{OpenAIFunction, OpenAIToolCall},
query_builder::{ParsedResponse, StreamEventProcessor},
types::StreamingEvent,
types::{ContentPart, OpenAIContent, OpenAIMessage, ToolDef},
};
use aws_config::BehaviorVersion;
use aws_credential_types::provider::token::ProvideToken;
use aws_sdk_bedrockruntime::types::{
ContentBlock, ConversationRole, ConverseStreamOutput, ImageBlock, ImageFormat, ImageSource,
InferenceConfiguration, Message, SystemContentBlock, Tool, ToolInputSchema, ToolSpecification,
};
use aws_sdk_bedrockruntime::Client as BedrockRuntimeClient;
use std::collections::HashMap;
use windmill_common::{client::AuthedClient, error::Error};
/// Constants for commonly used strings to avoid allocations
const FUNCTION_TYPE: &str = "function";
const EMPTY_JSON: &str = "{}";
#[derive(Debug, Clone)]
pub struct BearerTokenProvider {
token: String,
}
impl BearerTokenProvider {
pub fn new(token: String) -> Self {
Self { token }
}
}
impl ProvideToken for BearerTokenProvider {
fn provide_token<'a>(&'a self) -> aws_credential_types::provider::future::ProvideToken<'a>
where
Self: 'a,
{
aws_credential_types::provider::future::ProvideToken::ready(Ok(
aws_credential_types::Token::new(self.token.clone(), None),
))
}
}
pub struct BedrockClient {
client: BedrockRuntimeClient,
}
impl BedrockClient {
pub async fn from_bearer_token(bearer_token: String, region: &str) -> Result<Self, Error> {
let config = aws_sdk_bedrockruntime::config::Builder::new()
.region(aws_config::Region::new(region.to_string()))
.behavior_version(BehaviorVersion::latest())
.token_provider(BearerTokenProvider::new(bearer_token))
.build();
Ok(Self { client: BedrockRuntimeClient::from_conf(config) })
}
pub fn client(&self) -> &BedrockRuntimeClient {
&self.client
}
}
/// Format AWS SDK errors with detailed information
fn format_bedrock_error<E, R>(error: &aws_sdk_bedrockruntime::error::SdkError<E, R>) -> String
where
E: std::fmt::Debug + std::fmt::Display,
R: std::fmt::Debug,
{
use aws_sdk_bedrockruntime::error::SdkError;
match error {
SdkError::ServiceError(err) => {
// Include both the display and debug representations for maximum detail
format!("Service error: {} (details: {:?})", err.err(), err)
}
SdkError::ConstructionFailure(err) => {
format!("Request construction failed: {:?}", err)
}
SdkError::DispatchFailure(err) => {
format!("Request dispatch failed: {:?}", err)
}
SdkError::ResponseError(err) => {
format!("Response error: {:?}", err)
}
SdkError::TimeoutError(err) => {
format!("Request timeout: {:?}", err)
}
_ => format!("{:?}", error),
}
}
/// Convert AWS Smithy Document to serde_json::Value
fn document_to_json(doc: &aws_smithy_types::Document) -> serde_json::Value {
use aws_smithy_types::Document;
match doc {
Document::Object(map) => {
let mut obj = serde_json::Map::new();
for (k, v) in map {
obj.insert(k.clone(), document_to_json(v));
}
serde_json::Value::Object(obj)
}
Document::Array(arr) => {
serde_json::Value::Array(arr.iter().map(document_to_json).collect())
}
Document::Number(num) => {
// Try to parse as different number types
serde_json::Value::Number(
serde_json::Number::from_f64(num.to_f64_lossy())
.unwrap_or(serde_json::Number::from(0)),
)
}
Document::String(s) => serde_json::Value::String(s.clone()),
Document::Bool(b) => serde_json::Value::Bool(*b),
Document::Null => serde_json::Value::Null,
}
}
/// Convert serde_json::Value to AWS Smithy Document
fn json_to_document(value: serde_json::Value) -> aws_smithy_types::Document {
use aws_smithy_types::Document;
use serde_json::Value;
match value {
Value::Object(map) => {
let mut doc_map = std::collections::HashMap::new();
for (k, v) in map {
doc_map.insert(k, json_to_document(v));
}
Document::Object(doc_map)
}
Value::Array(arr) => Document::Array(arr.into_iter().map(json_to_document).collect()),
Value::Number(num) => {
if let Some(i) = num.as_i64() {
Document::Number(aws_smithy_types::Number::PosInt(i as u64))
} else if let Some(f) = num.as_f64() {
Document::Number(aws_smithy_types::Number::Float(f))
} else {
Document::Number(aws_smithy_types::Number::PosInt(0))
}
}
Value::String(s) => Document::String(s),
Value::Bool(b) => Document::Bool(b),
Value::Null => Document::Null,
}
}
/// Convert OpenAI-style messages to Bedrock format
///
/// Separates system messages from conversation messages as required by Bedrock API.
///
/// # Returns
/// Tuple of (conversation_messages, system_prompts)
pub fn openai_messages_to_bedrock(
messages: &[OpenAIMessage],
) -> Result<(Vec<Message>, Vec<SystemContentBlock>), Error> {
let mut bedrock_messages = Vec::new();
let mut system_prompts = Vec::new();
for msg in messages {
match msg.role.as_str() {
"system" => {
// Extract system messages separately
if let Some(ref content) = msg.content {
let text = content_to_text(content);
if !text.is_empty() {
system_prompts.push(SystemContentBlock::Text(text));
}
}
}
"user" | "assistant" => {
bedrock_messages.push(convert_message(msg)?);
}
"tool" => {
// Tool results are handled as user messages with ToolResult content
bedrock_messages.push(convert_tool_message(msg)?);
}
_ => {
return Err(Error::BadRequest(format!("Unsupported role: {}", msg.role)));
}
}
}
Ok((bedrock_messages, system_prompts))
}
/// Helper to extract text from OpenAIContent (ignoring images)
fn content_to_text(content: &OpenAIContent) -> String {
match content {
OpenAIContent::Text(text) => text.to_string(),
OpenAIContent::Parts(parts) => {
// Extract only text parts and join them
let text_parts: Vec<&str> = parts
.iter()
.filter_map(|part| match part {
ContentPart::Text { text } => Some(text.as_str()),
_ => None,
})
.collect();
text_parts.join(" ")
}
}
}
/// Parse image data URL and extract format and base64 data
fn parse_image_data_url(url: &str) -> Result<(ImageFormat, Vec<u8>), Error> {
if !url.starts_with("data:") {
return Err(Error::internal_err("Image URL must be a data URL"));
}
// Parse data:image/png;base64,<data>
let base64_start = url
.find("base64,")
.ok_or_else(|| Error::internal_err("Invalid data URL format"))?;
let base64_data = &url[base64_start + 7..];
let mime_type = url
.split(';')
.next()
.and_then(|s| s.strip_prefix("data:"))
.unwrap_or("image/png");
// Extract format from MIME type (e.g., "image/png" -> "png")
let format_str = mime_type
.rsplit_once('/')
.map(|(_, format)| format)
.unwrap_or("png");
// Map to ImageFormat enum
let format = match format_str {
"png" => ImageFormat::Png,
"jpeg" | "jpg" => ImageFormat::Jpeg,
"gif" => ImageFormat::Gif,
"webp" => ImageFormat::Webp,
_ => ImageFormat::Png, // Default to PNG
};
// Decode base64
let bytes = base64::Engine::decode(&base64::engine::general_purpose::STANDARD, base64_data)
.map_err(|e| Error::internal_err(format!("Failed to decode base64 image: {}", e)))?;
Ok((format, bytes))
}
/// Convert a ContentPart to Bedrock ContentBlock
fn content_part_to_block(part: &ContentPart) -> Result<Option<ContentBlock>, Error> {
match part {
ContentPart::Text { text } => {
if text.is_empty() {
Ok(None)
} else {
Ok(Some(ContentBlock::Text(text.clone())))
}
}
ContentPart::ImageUrl { image_url } => {
let (format, bytes) = parse_image_data_url(&image_url.url)?;
let image_source = ImageSource::Bytes(bytes.into());
let image_block = ImageBlock::builder()
.format(format)
.source(image_source)
.build()
.map_err(|e| Error::internal_err(format!("Failed to build image block: {}", e)))?;
Ok(Some(ContentBlock::Image(image_block)))
}
ContentPart::S3Object { .. } => {
// S3Objects are already converted to ImageUrl by prepare_messages_for_api
// If we somehow get here, skip it
Ok(None)
}
}
}
/// Convert a single OpenAI message to Bedrock Message
fn convert_message(msg: &OpenAIMessage) -> Result<Message, Error> {
let role = match msg.role.as_str() {
"user" => ConversationRole::User,
"assistant" => ConversationRole::Assistant,
_ => {
return Err(Error::internal_err(format!(
"Unsupported role: {}",
msg.role
)));
}
};
let mut content_blocks = Vec::new();
// Handle content (text and/or images)
if let Some(ref content) = msg.content {
match content {
OpenAIContent::Text(text) => {
if !text.is_empty() {
content_blocks.push(ContentBlock::Text(text.clone()));
}
}
OpenAIContent::Parts(parts) => {
for part in parts {
if let Some(block) = content_part_to_block(part)? {
content_blocks.push(block);
}
}
}
}
}
// Handle tool calls (for assistant messages)
if let Some(ref tool_calls) = msg.tool_calls {
for tc in tool_calls {
content_blocks.push(convert_tool_call_to_content(tc)?);
}
}
// Bedrock requires at least one content block
if content_blocks.is_empty() {
content_blocks.push(ContentBlock::Text(String::new()));
}
Message::builder()
.role(role)
.set_content(Some(content_blocks))
.build()
.map_err(|e| Error::internal_err(format!("Failed to build message: {}", e)))
}
/// Convert OpenAI tool call to Bedrock ToolUse content block
fn convert_tool_call_to_content(tool_call: &OpenAIToolCall) -> Result<ContentBlock, Error> {
let input = json_to_document(
serde_json::from_str(&tool_call.function.arguments)
.unwrap_or_else(|_| serde_json::json!({})),
);
Ok(ContentBlock::ToolUse(
aws_sdk_bedrockruntime::types::ToolUseBlock::builder()
.tool_use_id(&tool_call.id)
.name(&tool_call.function.name)
.input(input)
.build()
.map_err(|e| Error::internal_err(format!("Failed to build tool use: {}", e)))?,
))
}
/// Convert tool result message to Bedrock format
fn convert_tool_message(msg: &OpenAIMessage) -> Result<Message, Error> {
let tool_call_id = msg
.tool_call_id
.as_ref()
.ok_or_else(|| Error::internal_err("Tool message missing tool_call_id"))?;
let content_str = msg
.content
.as_ref()
.map(|c| content_to_text(c))
.unwrap_or_default();
// Try to parse as JSON, otherwise use text
let tool_result_content =
if let Ok(json_val) = serde_json::from_str::<serde_json::Value>(&content_str) {
if json_val.is_object() {
vec![aws_sdk_bedrockruntime::types::ToolResultContentBlock::Json(
json_to_document(json_val),
)]
} else {
// Wrap primitives and arrays in an object
vec![aws_sdk_bedrockruntime::types::ToolResultContentBlock::Json(
json_to_document(serde_json::json!({"result": json_val})),
)]
}
} else {
vec![aws_sdk_bedrockruntime::types::ToolResultContentBlock::Text(
content_str.to_string(),
)]
};
let tool_result = ContentBlock::ToolResult(
aws_sdk_bedrockruntime::types::ToolResultBlock::builder()
.tool_use_id(tool_call_id)
.set_content(Some(tool_result_content))
.build()
.map_err(|e| Error::internal_err(format!("Failed to build tool result: {}", e)))?,
);
Message::builder()
.role(ConversationRole::User)
.content(tool_result)
.build()
.map_err(|e| Error::internal_err(format!("Failed to build tool result message: {}", e)))
}
/// Convert OpenAI tool definitions to Bedrock format
pub fn openai_tools_to_bedrock(tools: &[ToolDef]) -> Result<Vec<Tool>, Error> {
tools
.iter()
.map(|tool_def| {
let spec = &tool_def.function;
// Convert parameters (RawValue) to Document via serde_json::Value
let param_value: serde_json::Value = serde_json::from_str(spec.parameters.get())
.map_err(|e| Error::internal_err(format!("Invalid tool schema: {}", e)))?;
let input_schema = ToolInputSchema::Json(json_to_document(param_value));
let tool_spec = ToolSpecification::builder()
.name(&spec.name)
.set_description(spec.description.clone())
.input_schema(input_schema)
.build()
.map_err(|e| Error::internal_err(format!("Failed to build tool spec: {}", e)))?;
Ok(Tool::ToolSpec(tool_spec))
})
.collect()
}
/// Create inference configuration from parameters
pub fn create_inference_config(
temperature: Option<f32>,
max_tokens: Option<i32>,
) -> Option<InferenceConfiguration> {
if temperature.is_none() && max_tokens.is_none() {
return None;
}
let mut builder = InferenceConfiguration::builder();
if let Some(temp) = temperature {
builder = builder.temperature(temp);
}
if let Some(max_tok) = max_tokens {
builder = builder.max_tokens(max_tok);
}
Some(builder.build())
}
/// Extract text content and tool calls from Bedrock Converse response
pub fn bedrock_response_to_openai(
output: &aws_sdk_bedrockruntime::operation::converse::ConverseOutput,
) -> Result<(Option<String>, Vec<OpenAIToolCall>), Error> {
let mut text_content = String::new();
let mut tool_calls = Vec::new();
if let Some(message) = output.output().and_then(|o| o.as_message().ok()) {
let content_blocks = message.content();
if !content_blocks.is_empty() {
for block in content_blocks {
match block {
ContentBlock::Text(text) => {
text_content.push_str(&text);
}
ContentBlock::ToolUse(tool_use) => {
// Convert to OpenAI tool call format
// Convert aws_smithy_types::Document to serde_json::Value
let input_value = document_to_json(tool_use.input());
let arguments = serde_json::to_string(&input_value)
.unwrap_or_else(|_| EMPTY_JSON.to_string());
tool_calls.push(OpenAIToolCall {
id: tool_use.tool_use_id().to_string(),
r#type: FUNCTION_TYPE.to_string(),
function: OpenAIFunction {
name: tool_use.name().to_string(),
arguments,
},
});
}
_ => {}
}
}
}
}
let content = if text_content.is_empty() {
None
} else {
Some(text_content)
};
Ok((content, tool_calls))
}
/// Extract text delta from Bedrock stream event
pub fn bedrock_stream_event_to_text(event: &ConverseStreamOutput) -> Option<String> {
match event {
ConverseStreamOutput::ContentBlockDelta(delta) => delta
.delta()
.and_then(|d| d.as_text().ok())
.map(|s| s.to_string()),
_ => None,
}
}
/// Represents a streaming tool call being accumulated
#[derive(Debug, Clone)]
pub struct StreamingToolCall {
pub id: String,
pub name: String,
pub arguments: String,
}
/// Extract tool use start event from stream
pub fn bedrock_stream_event_to_tool_start(
event: &ConverseStreamOutput,
) -> Option<StreamingToolCall> {
match event {
ConverseStreamOutput::ContentBlockStart(start) => {
if let Some(tool_use) = start.start().and_then(|s| s.as_tool_use().ok()) {
Some(StreamingToolCall {
id: tool_use.tool_use_id().to_string(),
name: tool_use.name().to_string(),
arguments: String::new(),
})
} else {
None
}
}
_ => None,
}
}
/// Extract tool use input delta from stream
pub fn bedrock_stream_event_to_tool_delta(event: &ConverseStreamOutput) -> Option<String> {
match event {
ConverseStreamOutput::ContentBlockDelta(delta) => delta
.delta()
.and_then(|d| d.as_tool_use().ok())
.map(|tool_use| tool_use.input().to_string()),
_ => None,
}
}
/// Check if stream event indicates content block stop
pub fn bedrock_stream_event_is_block_stop(event: &ConverseStreamOutput) -> bool {
matches!(event, ConverseStreamOutput::ContentBlockStop(_))
}
/// Convert accumulated streaming tool calls to OpenAI format
pub fn streaming_tool_calls_to_openai(tool_calls: Vec<StreamingToolCall>) -> Vec<OpenAIToolCall> {
tool_calls
.into_iter()
.map(|tc| OpenAIToolCall {
id: tc.id,
function: OpenAIFunction { name: tc.name, arguments: tc.arguments },
r#type: FUNCTION_TYPE.to_string(),
})
.collect()
}
#[derive(Default)]
pub struct BedrockQueryBuilder;
impl BedrockQueryBuilder {
/// Execute Bedrock request (streaming or non-streaming)
pub async fn execute_request(
&self,
messages: &[OpenAIMessage],
tools: Option<&[ToolDef]>,
model: &str,
temperature: Option<f32>,
max_tokens: Option<u32>,
api_key: &str,
region: &str,
should_stream: bool,
stream_event_processor: Option<StreamEventProcessor>,
client: &AuthedClient,
workspace_id: &str,
structured_output_tool_name: Option<&str>,
) -> Result<ParsedResponse, Error> {
// Create Bedrock client with bearer token authentication
let bedrock_client = BedrockClient::from_bearer_token(api_key.to_string(), region).await?;
// Prepare messages: convert S3Objects to ImageUrls by downloading from S3
let prepared_messages = prepare_messages_for_api(messages, client, workspace_id).await?;
// Convert messages to Bedrock format (separates system prompts)
let (bedrock_messages, system_prompts) = openai_messages_to_bedrock(&prepared_messages)?;
// Build inference configuration
let inference_config = create_inference_config(temperature, max_tokens.map(|t| t as i32));
// Build tool configuration with optional ToolChoice
let tool_config = self.build_tool_config(tools, structured_output_tool_name.is_some())?;
if should_stream {
self.execute_converse_stream(
&bedrock_client,
model,
bedrock_messages,
system_prompts,
inference_config,
tool_config,
stream_event_processor,
)
.await
} else {
self.execute_converse(
&bedrock_client,
model,
bedrock_messages,
system_prompts,
inference_config,
tool_config,
)
.await
}
}
/// Build tool configuration with optional ToolChoice for structured output
fn build_tool_config(
&self,
tools: Option<&[ToolDef]>,
force_tool_use: bool,
) -> Result<Option<aws_sdk_bedrockruntime::types::ToolConfiguration>, Error> {
if let Some(tools) = tools {
let bedrock_tools = openai_tools_to_bedrock(tools)?;
let mut tool_config_builder =
aws_sdk_bedrockruntime::types::ToolConfiguration::builder()
.set_tools(Some(bedrock_tools));
// For structured output, force the model to use the tool
if force_tool_use {
tool_config_builder = tool_config_builder.tool_choice(
aws_sdk_bedrockruntime::types::ToolChoice::Any(
aws_sdk_bedrockruntime::types::AnyToolChoice::builder().build(),
),
);
}
Ok(Some(tool_config_builder.build().map_err(|e| {
Error::internal_err(format!("Failed to build tool configuration: {}", e))
})?))
} else {
Ok(None)
}
}
/// Execute non-streaming Bedrock request
async fn execute_converse(
&self,
bedrock_client: &BedrockClient,
model: &str,
bedrock_messages: Vec<aws_sdk_bedrockruntime::types::Message>,
system_prompts: Vec<aws_sdk_bedrockruntime::types::SystemContentBlock>,
inference_config: Option<aws_sdk_bedrockruntime::types::InferenceConfiguration>,
tool_config: Option<aws_sdk_bedrockruntime::types::ToolConfiguration>,
) -> Result<ParsedResponse, Error> {
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));
}
// Execute the request
let response = request_builder.send().await.map_err(|e| {
let error_msg = format!("Bedrock API error: {}", format_bedrock_error(&e));
Error::internal_err(error_msg)
})?;
// Convert response back to OpenAI format
let (content, tool_calls) = bedrock_response_to_openai(&response)?;
Ok(ParsedResponse::Text { content, tool_calls, events_str: None })
}
/// Execute streaming Bedrock request
async fn execute_converse_stream(
&self,
bedrock_client: &BedrockClient,
model: &str,
bedrock_messages: Vec<aws_sdk_bedrockruntime::types::Message>,
system_prompts: Vec<aws_sdk_bedrockruntime::types::SystemContentBlock>,
inference_config: Option<aws_sdk_bedrockruntime::types::InferenceConfiguration>,
tool_config: Option<aws_sdk_bedrockruntime::types::ToolConfiguration>,
stream_event_processor: Option<StreamEventProcessor>,
) -> Result<ParsedResponse, Error> {
// Build streaming 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));
}
// Execute streaming request
let mut stream = request_builder
.send()
.await
.map_err(|e| {
let error_msg =
format!("Bedrock streaming API error: {}", format_bedrock_error(&e));
Error::internal_err(error_msg)
})?
.stream;
let mut accumulated_text = String::new();
let mut events_str = String::new();
let mut accumulated_tool_calls: HashMap<String, StreamingToolCall> = HashMap::new();
let mut current_tool_use_id: Option<String> = None;
// Process stream events
loop {
match stream.recv().await {
Ok(Some(event)) => {
// Handle tool use start
if let Some(tool_call) = bedrock_stream_event_to_tool_start(&event) {
current_tool_use_id = Some(tool_call.id.clone());
accumulated_tool_calls.insert(tool_call.id.clone(), tool_call);
}
// Handle text delta
if let Some(text_delta) = bedrock_stream_event_to_text(&event) {
accumulated_text.push_str(&text_delta);
if let Some(processor) = stream_event_processor.as_ref() {
processor
.send(
StreamingEvent::TokenDelta { content: text_delta },
&mut events_str,
)
.await?;
}
}
// Handle tool use input delta
if let Some(input_delta) = bedrock_stream_event_to_tool_delta(&event) {
if let Some(tool_id) = &current_tool_use_id {
if let Some(tool_call) = accumulated_tool_calls.get_mut(tool_id) {
tool_call.arguments.push_str(&input_delta);
}
}
}
// Handle content block stop
if bedrock_stream_event_is_block_stop(&event) {
current_tool_use_id = None;
}
}
Ok(None) => break, // Stream ended
Err(e) => {
return Err(Error::internal_err(format!("Bedrock stream error: {}", e)));
}
}
}
// Send tool call events to stream processor
if let Some(processor) = stream_event_processor.as_ref() {
for tool_call in accumulated_tool_calls.values() {
processor
.send(
StreamingEvent::ToolCallArguments {
call_id: tool_call.id.clone(),
function_name: tool_call.name.clone(),
arguments: tool_call.arguments.clone(),
},
&mut events_str,
)
.await?;
}
}
let content = if accumulated_text.is_empty() {
None
} else {
Some(accumulated_text)
};
let tool_calls =
streaming_tool_calls_to_openai(accumulated_tool_calls.into_values().collect());
Ok(ParsedResponse::Text {
content,
tool_calls,
events_str: if events_str.is_empty() {
None
} else {
Some(events_str)
},
})
}
}
@@ -255,7 +255,13 @@ impl QueryBuilder for GoogleAIQueryBuilder {
.await
}
fn get_endpoint(&self, base_url: &str, model: &str, output_type: &OutputType) -> String {
fn get_endpoint(
&self,
base_url: &str,
model: &str,
output_type: &OutputType,
_stream: bool,
) -> String {
match output_type {
OutputType::Text => format!("{}/chat/completions", base_url), // Use OpenAI-compatible endpoint
OutputType::Image => {
@@ -1,3 +1,4 @@
pub mod bedrock;
pub mod google_ai;
pub mod openai;
pub mod openrouter;
@@ -4,14 +4,13 @@ use serde_json;
use windmill_common::{ai_providers::AIProvider, client::AuthedClient, error::Error};
use crate::ai::{
image_handler::download_and_encode_s3_image,
image_handler::{download_and_encode_s3_image, prepare_messages_for_api},
query_builder::{BuildRequestArgs, ParsedResponse, QueryBuilder, StreamEventProcessor},
sse::{OpenAISSEParser, SSEParser},
types::*,
utils::is_claude_model,
utils::should_use_structured_output_tool,
};
// OpenAI-specific types
#[derive(Deserialize, Serialize, Clone, Debug)]
pub struct OpenAIFunction {
pub name: String,
@@ -114,62 +113,6 @@ impl OpenAIQueryBuilder {
Self { provider_kind }
}
pub async fn prepare_messages_for_api(
&self,
messages: &[OpenAIMessage],
client: &AuthedClient,
workspace_id: &str,
) -> Result<Vec<OpenAIMessage>, Error> {
let mut prepared_messages = Vec::new();
for message in messages {
let mut prepared_message = message.clone();
if let Some(content) = &message.content {
match content {
OpenAIContent::Text(text) => {
prepared_message.content = Some(OpenAIContent::Text(text.clone()));
}
OpenAIContent::Parts(parts) => {
let mut prepared_content = Vec::new();
for part in parts {
match part {
ContentPart::S3Object { s3_object } => {
// Convert S3Object to base64 image URL
let (mime_type, image_bytes) = download_and_encode_s3_image(
s3_object,
client,
workspace_id,
)
.await?;
prepared_content.push(ContentPart::ImageUrl {
image_url: ImageUrlData {
url: format!(
"data:{};base64,{}",
mime_type, image_bytes
),
},
});
}
other => {
// Keep Text and ImageUrl as-is
prepared_content.push(other.clone());
}
}
}
prepared_message.content = Some(OpenAIContent::Parts(prepared_content));
}
}
}
prepared_messages.push(prepared_message);
}
Ok(prepared_messages)
}
async fn build_text_request(
&self,
args: &BuildRequestArgs<'_>,
@@ -177,9 +120,8 @@ impl OpenAIQueryBuilder {
workspace_id: &str,
stream: bool,
) -> Result<String, Error> {
let prepared_messages = self
.prepare_messages_for_api(args.messages, client, workspace_id)
.await?;
let prepared_messages =
prepare_messages_for_api(args.messages, client, workspace_id).await?;
// Check if we need to add response_format for structured output
let has_output_properties = args
@@ -203,9 +145,10 @@ impl OpenAIQueryBuilder {
None
};
let is_claude_model = is_claude_model(&args.model);
let should_use_structured_output_tool =
should_use_structured_output_tool(&self.provider_kind, args.model);
// Force usage of structured output tool for Claude models when structured output provided
let tool_choice = if is_claude_model && response_format.is_some() {
let tool_choice = if should_use_structured_output_tool && response_format.is_some() {
Some(ToolChoice::Required)
} else {
None
@@ -405,7 +348,13 @@ impl QueryBuilder for OpenAIQueryBuilder {
})
}
fn get_endpoint(&self, base_url: &str, model: &str, output_type: &OutputType) -> String {
fn get_endpoint(
&self,
base_url: &str,
model: &str,
output_type: &OutputType,
_stream: bool,
) -> String {
let path = match output_type {
OutputType::Text => "chat/completions",
OutputType::Image => "responses",
@@ -4,6 +4,7 @@ use serde_json;
use windmill_common::{ai_providers::AIProvider, client::AuthedClient, error::Error};
use crate::ai::{
image_handler::prepare_messages_for_api,
providers::openai::{OpenAIQueryBuilder, OpenAIResponse},
query_builder::{BuildRequestArgs, ParsedResponse, QueryBuilder, StreamEventProcessor},
types::*,
@@ -91,11 +92,9 @@ impl QueryBuilder for OpenRouterQueryBuilder {
}
OutputType::Image => {
// For image generation, we need to add modalities field
// First, prepare the messages using the OpenAI builder's logic
let openai_builder = &self.openai_builder;
let prepared_messages = openai_builder
.prepare_messages_for_api(args.messages, client, workspace_id)
.await?;
// First, prepare the messages
let prepared_messages =
prepare_messages_for_api(args.messages, client, workspace_id).await?;
// Check if we need to add response_format for structured output
let has_output_properties = args
@@ -204,7 +203,13 @@ impl QueryBuilder for OpenRouterQueryBuilder {
.await
}
fn get_endpoint(&self, base_url: &str, _model: &str, _output_type: &OutputType) -> String {
fn get_endpoint(
&self,
base_url: &str,
_model: &str,
_output_type: &OutputType,
_stream: bool,
) -> String {
// OpenRouter uses the same endpoint for both text and image generation
format!("{}/chat/completions", base_url)
}
@@ -69,7 +69,13 @@ pub trait QueryBuilder: Send + Sync {
}
/// Get the API endpoint for this provider
fn get_endpoint(&self, base_url: &str, model: &str, output_type: &OutputType) -> String;
fn get_endpoint(
&self,
base_url: &str,
model: &str,
output_type: &OutputType,
stream: bool,
) -> String;
/// Get the authentication headers for this provider
fn get_auth_headers(
+10 -1
View File
@@ -126,6 +126,7 @@ pub struct ProviderResource {
pub api_key: String,
#[serde(alias = "baseUrl")]
pub base_url: Option<String>,
pub region: Option<String>,
}
#[derive(Deserialize, Debug)]
@@ -146,9 +147,17 @@ impl ProviderWithResource {
pub async fn get_base_url(&self, db: &DB) -> Result<String, Error> {
self.kind
.get_base_url(self.resource.base_url.clone(), db)
.get_base_url(
self.resource.base_url.clone(),
self.resource.region.clone(),
db,
)
.await
}
pub fn get_region(&self) -> Option<&str> {
self.resource.region.as_deref()
}
}
#[derive(Serialize)]
+4 -3
View File
@@ -8,6 +8,7 @@ use std::{
};
use uuid::Uuid;
use windmill_common::{
ai_providers::AIProvider,
db::DB,
error::Error,
flow_conversations::{add_message_to_conversation_tx, MessageType},
@@ -311,9 +312,9 @@ pub fn get_step_name_from_flow(
)
}
/// Claude models starts with claude if provider is anthropic, or anthropic for openrouter and other providers
pub fn is_claude_model(model: &str) -> bool {
model.starts_with("claude") || model.starts_with("anthropic")
/// AWS Bedrock do not handle structured output query param, so we use a tool for structured output. Same for every Claude models.
pub fn should_use_structured_output_tool(provider: &AIProvider, model: &str) -> bool {
model.contains("claude") || provider == &AIProvider::AWSBedrock
}
/// Cleanup MCP clients by gracefully shutting down connections
+279 -241
View File
@@ -2,9 +2,9 @@ use crate::ai::tools::{execute_tool_calls, ToolExecutionContext};
use crate::ai::utils::{
add_message_to_conversation, any_tool_needs_previous_result, cleanup_mcp_clients,
filter_schema_by_input_transforms, find_unique_tool_name, get_flow_context,
get_flow_job_runnable_and_raw_flow, get_step_name_from_flow, is_claude_model, load_mcp_tools,
parse_raw_script_schema, update_flow_status_module_with_actions,
update_flow_status_module_with_actions_success,
get_flow_job_runnable_and_raw_flow, get_step_name_from_flow, load_mcp_tools,
parse_raw_script_schema, should_use_structured_output_tool,
update_flow_status_module_with_actions, update_flow_status_module_with_actions_success,
};
use crate::memory_oss::{read_from_memory, write_to_memory};
use crate::worker_flow::{get_previous_job_result, get_transform_context};
@@ -15,7 +15,7 @@ use std::{collections::HashMap, sync::Arc};
use uuid::Uuid;
use windmill_common::mcp_client::McpClient;
use windmill_common::{
ai_providers::AZURE_API_VERSION,
ai_providers::{AIProvider, AZURE_API_VERSION},
cache,
client::AuthedClient,
db::DB,
@@ -376,6 +376,7 @@ pub async fn run_agent(
let output_type = args.output_type.as_ref().unwrap_or(&OutputType::Text);
let base_url = args.provider.get_base_url(db).await?;
let api_key = args.provider.get_api_key();
let region = args.provider.get_region();
// Create the query builder for the provider
let query_builder = create_query_builder(&args.provider);
@@ -498,14 +499,15 @@ pub async fn run_agent(
.map(|props| !props.is_empty())
.unwrap_or(false);
let is_claude_model = is_claude_model(&args.provider.model);
let should_use_structured_output_tool =
should_use_structured_output_tool(&args.provider.kind, &args.provider.model);
let mut used_structured_output_tool = false;
let mut structured_output_tool_name: Option<String> = None;
// For text output with schema, handle structured output
if has_output_properties && output_type == &OutputType::Text {
let schema = args.output_schema.as_ref().unwrap();
if is_claude_model {
if should_use_structured_output_tool {
// Anthropic uses a tool for structured output
let unique_tool_name = find_unique_tool_name("structured_output", tool_defs.as_deref());
structured_output_tool_name = Some(unique_tool_name.clone());
@@ -551,253 +553,289 @@ pub async fn run_agent(
break;
}
// For text output or image output with tools
let build_args = BuildRequestArgs {
messages: &messages,
tools: tool_defs.as_deref(),
model: args.provider.get_model(),
temperature: args.temperature,
max_tokens: args.max_completion_tokens,
output_schema: args.output_schema.as_ref(),
output_type,
system_prompt: args.system_prompt.as_deref(),
user_message: &args.user_message,
images: args.user_images.as_deref(),
};
// Special handling for AWS Bedrock using the official SDK
let parsed = if args.provider.kind == AIProvider::AWSBedrock {
let Some(region) = region else {
return Err(Error::internal_err(
"AWS Bedrock region is required".to_string(),
));
};
// Use Bedrock SDK via dedicated query builder
crate::ai::providers::bedrock::BedrockQueryBuilder::default()
.execute_request(
&messages,
tool_defs.as_deref(),
args.provider.get_model(),
args.temperature,
args.max_completion_tokens,
api_key,
region,
should_stream,
stream_event_processor.clone(),
client,
&job.workspace_id,
structured_output_tool_name.as_deref(),
)
.await?
} else {
// For non-Bedrock providers, use HTTP client
let build_args = BuildRequestArgs {
messages: &messages,
tools: tool_defs.as_deref(),
model: args.provider.get_model(),
temperature: args.temperature,
max_tokens: args.max_completion_tokens,
output_schema: args.output_schema.as_ref(),
output_type,
system_prompt: args.system_prompt.as_deref(),
user_message: &args.user_message,
images: args.user_images.as_deref(),
};
let request_body = query_builder
.build_request(&build_args, client, &job.workspace_id, should_stream)
.await?;
let request_body = query_builder
.build_request(&build_args, client, &job.workspace_id, should_stream)
.await?;
let endpoint =
query_builder.get_endpoint(&base_url, args.provider.get_model(), output_type);
let auth_headers = query_builder.get_auth_headers(api_key, &base_url, output_type);
let endpoint = query_builder.get_endpoint(
&base_url,
args.provider.get_model(),
output_type,
should_stream,
);
let auth_headers = query_builder.get_auth_headers(api_key, &base_url, output_type);
let timeout = resolve_job_timeout(conn, &job.workspace_id, job.id, job.timeout)
.await
.0;
let timeout = resolve_job_timeout(conn, &job.workspace_id, job.id, job.timeout)
.await
.0;
let mut request = HTTP_CLIENT
.post(&endpoint)
.timeout(timeout)
.header("Content-Type", "application/json");
let mut request = HTTP_CLIENT
.post(&endpoint)
.timeout(timeout)
.header("Content-Type", "application/json");
// Apply authentication headers
for (header_name, header_value) in &auth_headers {
request = request.header(*header_name, header_value.clone());
}
// Apply authentication headers
for (header_name, header_value) in &auth_headers {
request = request.header(*header_name, header_value.clone());
}
// Apply custom headers from AI_HTTP_HEADERS environment variable
for (header_name, header_value) in AI_HTTP_HEADERS.iter() {
request = request.header(header_name.as_str(), header_value.as_str());
}
// Apply custom headers from AI_HTTP_HEADERS environment variable
for (header_name, header_value) in AI_HTTP_HEADERS.iter() {
request = request.header(header_name.as_str(), header_value.as_str());
}
if args.provider.kind.is_azure_openai(&base_url) {
request = request.query(&[("api-version", AZURE_API_VERSION)])
}
if args.provider.kind.is_azure_openai(&base_url) {
request = request.query(&[("api-version", AZURE_API_VERSION)])
}
let resp = request
.body(request_body)
.send()
.await
.map_err(|e| Error::internal_err(format!("Failed to call API: {}", e)))?;
let resp = request
.body(request_body)
.send()
.await
.map_err(|e| Error::internal_err(format!("Failed to call API: {}", e)))?;
match resp.error_for_status_ref() {
Ok(_) => {
let parsed = if let Some(stream_event_processor) = stream_event_processor.clone() {
query_builder
.parse_streaming_response(resp, stream_event_processor)
.await?
} else {
// Handle non-streaming response
query_builder.parse_response(resp).await?
};
match parsed {
ParsedResponse::Text { content: response_content, tool_calls, events_str } => {
if let Some(events_str) = events_str {
final_events_str.push_str(&events_str);
}
if let Some(ref response_content) = response_content {
actions.push(AgentAction::Message {});
messages.push(OpenAIMessage {
role: "assistant".to_string(),
content: Some(OpenAIContent::Text(response_content.clone())),
agent_action: Some(AgentAction::Message {}),
..Default::default()
});
update_flow_status_module_with_actions(db, parent_job, &actions)
.await?;
update_flow_status_module_with_actions_success(db, parent_job, true)
.await?;
content = Some(OpenAIContent::Text(response_content.clone()));
// Add assistant message to conversation if chat_input_enabled
let chat_enabled = flow_context
.flow_status
.as_ref()
.and_then(|fs| fs.chat_input_enabled)
.unwrap_or(false);
if chat_enabled && !response_content.is_empty() {
if let Some(memory_id) = flow_context
.flow_status
.as_ref()
.and_then(|fs| fs.memory_id)
{
let agent_job_id = job.id;
let db_clone = db.clone();
let message_content = response_content.clone();
let step_name = get_step_name_from_flow(
summary.as_deref(),
job.flow_step_id.as_deref(),
);
// Spawn task because we do not need to wait for the result
tokio::spawn(async move {
if let Err(e) = add_message_to_conversation(
&db_clone,
&memory_id,
Some(agent_job_id),
&message_content,
MessageType::Assistant,
&step_name,
true,
)
.await
{
tracing::warn!("Failed to add assistant message to conversation {}: {}", memory_id, e);
}
});
}
}
}
if tool_calls.is_empty() {
break;
} else if i == MAX_AGENT_ITERATIONS - 1 {
return Err(Error::internal_err(
"AI agent reached max iterations, but there are still tool calls"
.to_string(),
));
}
messages.push(OpenAIMessage {
role: "assistant".to_string(),
tool_calls: Some(tool_calls.clone()),
..Default::default()
});
// Handle tool calls using extracted tools module
let tool_execution_ctx = ToolExecutionContext {
db,
conn,
job,
parent_job,
summary: &summary,
client,
worker_dir,
base_internal_url,
worker_name,
hostname,
occupancy_metrics,
job_completed_tx,
killpill_rx,
stream_event_processor: stream_event_processor.as_ref(),
flow_context: &mut flow_context,
previous_result: &previous_result,
id_context: &id_context,
};
let (tool_messages, tool_content, tool_used_structured_output) =
execute_tool_calls(
tool_execution_ctx,
&tool_calls,
&tools,
mcp_clients,
&mut actions,
&mut final_events_str,
&structured_output_tool_name,
)
.await?;
messages.extend(tool_messages);
if let Some(tc) = tool_content {
content = Some(tc);
}
used_structured_output_tool = tool_used_structured_output;
}
ParsedResponse::Image { base64_data } => {
// For image output, upload to S3 and track in conversation
let s3_object = upload_image_to_s3(&base64_data, job, client).await?;
let content = to_raw_value(&s3_object);
// Add assistant message to conversation if chat_input_enabled
let chat_enabled = flow_context
.flow_status
.as_ref()
.and_then(|fs| fs.chat_input_enabled)
.unwrap_or(false);
if chat_enabled {
if let Some(memory_id) = flow_context
.flow_status
.as_ref()
.and_then(|fs| fs.memory_id)
{
let agent_job_id = job.id;
let db_clone = db.clone();
let flow_step_id_owned = job.flow_step_id.clone();
let summary_owned = summary.map(|s| s.to_string());
// Create extended version with type discriminator for conversation storage
// This avoids conflicts with outputs that are of the same format as S3 objects
let s3_with_type = S3ObjectWithType {
s3_object: s3_object.clone(),
r#type: "windmill_s3_object".to_string(),
};
let message_content = serde_json::to_string(&s3_with_type)
.unwrap_or_else(|_| content.get().to_string());
// Spawn task because we do not need to wait for the result
tokio::spawn(async move {
let step_name = get_step_name_from_flow(
summary_owned.as_deref(),
flow_step_id_owned.as_deref(),
);
if let Err(e) = add_message_to_conversation(
&db_clone,
&memory_id,
Some(agent_job_id),
&message_content,
MessageType::Assistant,
&step_name,
true,
)
.await
{
tracing::warn!("Failed to add assistant message to conversation {}: {}", memory_id, e);
}
});
}
}
// Return early since image generation is complete
return Ok(content);
match resp.error_for_status_ref() {
Ok(_) => {
if let Some(stream_event_processor) = stream_event_processor.clone() {
query_builder
.parse_streaming_response(resp, stream_event_processor)
.await?
} else {
// Handle non-streaming response
query_builder.parse_response(resp).await?
}
}
Err(e) => {
let _status = resp.status();
let text = resp
.text()
.await
.unwrap_or_else(|_| "<failed to read body>".to_string());
return Err(Error::internal_err(format!("API error: {} - {}", e, text)));
}
}
Err(e) => {
let _status = resp.status();
let text = resp
.text()
.await
.unwrap_or_else(|_| "<failed to read body>".to_string());
return Err(Error::internal_err(format!("API error: {} - {}", e, text)));
};
match parsed {
ParsedResponse::Text { content: response_content, tool_calls, events_str } => {
if let Some(events_str) = events_str {
final_events_str.push_str(&events_str);
}
if let Some(ref response_content) = response_content {
actions.push(AgentAction::Message {});
messages.push(OpenAIMessage {
role: "assistant".to_string(),
content: Some(OpenAIContent::Text(response_content.clone())),
agent_action: Some(AgentAction::Message {}),
..Default::default()
});
update_flow_status_module_with_actions(db, parent_job, &actions).await?;
update_flow_status_module_with_actions_success(db, parent_job, true).await?;
content = Some(OpenAIContent::Text(response_content.clone()));
// Add assistant message to conversation if chat_input_enabled
let chat_enabled = flow_context
.flow_status
.as_ref()
.and_then(|fs| fs.chat_input_enabled)
.unwrap_or(false);
if chat_enabled && !response_content.is_empty() {
if let Some(memory_id) = flow_context
.flow_status
.as_ref()
.and_then(|fs| fs.memory_id)
{
let agent_job_id = job.id;
let db_clone = db.clone();
let message_content = response_content.clone();
let step_name = get_step_name_from_flow(
summary.as_deref(),
job.flow_step_id.as_deref(),
);
// Spawn task because we do not need to wait for the result
tokio::spawn(async move {
if let Err(e) = add_message_to_conversation(
&db_clone,
&memory_id,
Some(agent_job_id),
&message_content,
MessageType::Assistant,
&step_name,
true,
)
.await
{
tracing::warn!(
"Failed to add assistant message to conversation {}: {}",
memory_id,
e
);
}
});
}
}
}
if tool_calls.is_empty() {
break;
} else if i == MAX_AGENT_ITERATIONS - 1 {
return Err(Error::internal_err(
"AI agent reached max iterations, but there are still tool calls"
.to_string(),
));
}
messages.push(OpenAIMessage {
role: "assistant".to_string(),
tool_calls: Some(tool_calls.clone()),
..Default::default()
});
// Handle tool calls using extracted tools module
let tool_execution_ctx = ToolExecutionContext {
db,
conn,
job,
parent_job,
summary: &summary,
client,
worker_dir,
base_internal_url,
worker_name,
hostname,
occupancy_metrics,
job_completed_tx,
killpill_rx,
stream_event_processor: stream_event_processor.as_ref(),
flow_context: &mut flow_context,
previous_result: &previous_result,
id_context: &id_context,
};
let (tool_messages, tool_content, tool_used_structured_output) =
execute_tool_calls(
tool_execution_ctx,
&tool_calls,
&tools,
mcp_clients,
&mut actions,
&mut final_events_str,
&structured_output_tool_name,
)
.await?;
messages.extend(tool_messages);
if let Some(tc) = tool_content {
content = Some(tc);
}
used_structured_output_tool = tool_used_structured_output;
}
ParsedResponse::Image { base64_data } => {
// For image output, upload to S3 and track in conversation
let s3_object = upload_image_to_s3(&base64_data, job, client).await?;
let content = to_raw_value(&s3_object);
// Add assistant message to conversation if chat_input_enabled
let chat_enabled = flow_context
.flow_status
.as_ref()
.and_then(|fs| fs.chat_input_enabled)
.unwrap_or(false);
if chat_enabled {
if let Some(memory_id) = flow_context
.flow_status
.as_ref()
.and_then(|fs| fs.memory_id)
{
let agent_job_id = job.id;
let db_clone = db.clone();
let flow_step_id_owned = job.flow_step_id.clone();
let summary_owned = summary.map(|s| s.to_string());
// Create extended version with type discriminator for conversation storage
// This avoids conflicts with outputs that are of the same format as S3 objects
let s3_with_type = S3ObjectWithType {
s3_object: s3_object.clone(),
r#type: "windmill_s3_object".to_string(),
};
let message_content = serde_json::to_string(&s3_with_type)
.unwrap_or_else(|_| content.get().to_string());
// Spawn task because we do not need to wait for the result
tokio::spawn(async move {
let step_name = get_step_name_from_flow(
summary_owned.as_deref(),
flow_step_id_owned.as_deref(),
);
if let Err(e) = add_message_to_conversation(
&db_clone,
&memory_id,
Some(agent_job_id),
&message_content,
MessageType::Assistant,
&step_name,
true,
)
.await
{
tracing::warn!(
"Failed to add assistant message to conversation {}: {}",
memory_id,
e
);
}
});
}
}
// Return early since image generation is complete
return Ok(content);
}
}
}
+96 -7
View File
@@ -76,6 +76,10 @@ export const AI_PROVIDERS: Record<AIProvider, AIProviderDetails> = {
label: 'Together AI',
defaultModels: ['meta-llama/Llama-3.3-70B-Instruct-Turbo']
},
aws_bedrock: {
label: 'AWS Bedrock',
defaultModels: ['global.anthropic.claude-haiku-4-5-20251001-v1:0']
},
customai: {
label: 'Custom AI',
defaultModels: []
@@ -100,18 +104,102 @@ export async function fetchAvailableModels(
provider: AIProvider,
signal?: AbortSignal
): Promise<string[]> {
const models = await fetch(`${location.origin}${OpenAPI.BASE}/w/${workspace}/ai/proxy/models`, {
signal,
headers: {
// Handle AWS Bedrock separately (needs both foundation-models and inference-profiles)
if (provider === 'aws_bedrock') {
const headers = {
'X-Resource-Path': resourcePath,
'X-Provider': provider,
...(provider === 'anthropic' ? { 'anthropic-version': '2023-06-01' } : {})
'X-Provider': provider
}
})
// Fetch both foundation models and inference profiles
const [foundationModelsResp, inferenceProfilesResp] = await Promise.all([
fetch(`${location.origin}${OpenAPI.BASE}/w/${workspace}/ai/proxy/foundation-models`, {
signal,
headers
}),
fetch(`${location.origin}${OpenAPI.BASE}/w/${workspace}/ai/proxy/inference-profiles`, {
signal,
headers
})
])
if (!foundationModelsResp.ok) {
console.error('Failed to fetch foundation models', foundationModelsResp)
throw new Error('Failed to fetch foundation models for AWS Bedrock')
}
const foundationModelsData = (await foundationModelsResp.json()) as {
modelSummaries: Array<{
modelId: string
modelArn: string
inputModalities: string[]
outputModalities: string[]
inferenceTypesSupported: string[]
}>
}
// Inference profiles fetch might fail in some regions/accounts
let inferenceProfiles: Array<{
inferenceProfileId: string
models: Array<{ modelArn: string }>
}> = []
if (inferenceProfilesResp.ok) {
const inferenceProfilesData = (await inferenceProfilesResp.json()) as {
inferenceProfileSummaries: Array<{
inferenceProfileId: string
models: Array<{ modelArn: string }>
}>
}
inferenceProfiles = inferenceProfilesData.inferenceProfileSummaries || []
} else {
console.warn('Failed to fetch inference profiles, will use direct model IDs only')
}
// Filter to TEXT-capable models
const textModels = foundationModelsData.modelSummaries.filter(
(m) => m.inputModalities?.includes('TEXT') && m.outputModalities?.includes('TEXT')
)
const onDemandModels = textModels
.filter(
(model) =>
model.inferenceTypesSupported?.includes('ON_DEMAND') &&
!model.inferenceTypesSupported?.includes('INFERENCE_PROFILE')
)
.map((model) => model.modelId)
const inferenceModels = inferenceProfiles.map((profile) => profile.inferenceProfileId)
const modelIds = [...onDemandModels, ...inferenceModels]
// Sort by default models
const defaultModels = AI_PROVIDERS[provider]?.defaultModels || []
return modelIds.sort((a, b) => {
const aInDefault = defaultModels.includes(a)
const bInDefault = defaultModels.includes(b)
if (aInDefault && !bInDefault) return -1
if (!aInDefault && bInDefault) return 1
return 0
})
}
// Standard provider handling
const endpoint = 'models'
const models = await fetch(
`${location.origin}${OpenAPI.BASE}/w/${workspace}/ai/proxy/${endpoint}`,
{
signal,
headers: {
'X-Resource-Path': resourcePath,
'X-Provider': provider,
...(provider === 'anthropic' ? { 'anthropic-version': '2023-06-01' } : {})
}
}
)
if (!models.ok) {
console.error('Failed to fetch models for provider', provider, models)
throw new Error(`Failed to fetch models for provider ${provider}`)
}
const data = (await models.json()) as { data: ModelResponse[] }
if (data.data.length > 0) {
const sortFunc = (provider: AIProvider) => (a: string, b: string) => {
@@ -271,7 +359,8 @@ export const PROVIDER_COMPLETION_CONFIG_MAP: Record<AIProvider, ChatCompletionCr
...DEFAULT_COMPLETION_CONFIG,
seed: undefined
},
anthropic: DEFAULT_COMPLETION_CONFIG
anthropic: DEFAULT_COMPLETION_CONFIG,
aws_bedrock: DEFAULT_COMPLETION_CONFIG
} as const
class WorkspacedAIClients {