From 79ac6312e87afa3646bddc0f7e66fc4367dbff7c Mon Sep 17 00:00:00 2001 From: centdix <40307056+centdix@users.noreply.github.com> Date: Mon, 17 Nov 2025 19:57:59 +0100 Subject: [PATCH] feat(ai): handle aws bedrock as provider (#7155) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 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 * 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 * 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 * 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 * 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 * cleaning * cleaning * cleaning * better error * cleaning * cleaning * rm * rename * apply region --------- Co-authored-by: Claude * fix default * no panic * no print * use utils file * cleaning --------- Co-authored-by: Claude --- backend/CLAUDE.md | 13 + backend/Cargo.lock | 90 +- backend/Cargo.toml | 5 +- backend/windmill-api/openapi.yaml | 59 +- backend/windmill-api/src/ai.rs | 101 ++- backend/windmill-api/src/bedrock.rs | 602 +++++++++++++ backend/windmill-api/src/lib.rs | 1 + backend/windmill-common/Cargo.toml | 2 + backend/windmill-common/src/ai_providers.rs | 13 +- backend/windmill-worker/Cargo.toml | 4 + .../windmill-worker/src/ai/image_handler.rs | 52 ++ .../src/ai/providers/bedrock.rs | 804 ++++++++++++++++++ .../src/ai/providers/google_ai.rs | 8 +- .../windmill-worker/src/ai/providers/mod.rs | 1 + .../src/ai/providers/openai.rs | 79 +- .../src/ai/providers/openrouter.rs | 17 +- .../windmill-worker/src/ai/query_builder.rs | 8 +- backend/windmill-worker/src/ai/types.rs | 11 +- backend/windmill-worker/src/ai/utils.rs | 7 +- backend/windmill-worker/src/ai_executor.rs | 520 +++++------ frontend/src/lib/components/copilot/lib.ts | 103 ++- 21 files changed, 2109 insertions(+), 391 deletions(-) create mode 100644 backend/windmill-api/src/bedrock.rs create mode 100644 backend/windmill-worker/src/ai/providers/bedrock.rs diff --git a/backend/CLAUDE.md b/backend/CLAUDE.md index 1fe509a0d6..78faa0bf77 100644 --- a/backend/CLAUDE.md +++ b/backend/CLAUDE.md @@ -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 diff --git a/backend/Cargo.lock b/backend/Cargo.lock index f402f4dc20..d55a6487c9 100644 --- a/backend/Cargo.lock +++ b/backend/Cargo.lock @@ -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", diff --git a/backend/Cargo.toml b/backend/Cargo.toml index 2192e17055..b196852f58 100644 --- a/backend/Cargo.toml +++ b/backend/Cargo.toml @@ -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" diff --git a/backend/windmill-api/openapi.yaml b/backend/windmill-api/openapi.yaml index 02423978cb..c8345e4b38 100644 --- a/backend/windmill-api/openapi.yaml +++ b/backend/windmill-api/openapi.yaml @@ -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 diff --git a/backend/windmill-api/src/ai.rs b/backend/windmill-api/src/ai.rs index e4d07cfdc6..4b18a67ec9 100644 --- a/backend/windmill-api/src/ai.rs +++ b/backend/windmill-api/src/ai.rs @@ -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, organization_id: Option, + region: Option, } #[derive(Deserialize, Debug)] @@ -98,7 +100,9 @@ impl AIRequestConfig { ) -> Result { 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))) + } } diff --git a/backend/windmill-api/src/bedrock.rs b/backend/windmill-api/src/bedrock.rs new file mode 100644 index 0000000000..ea93b5b3ea --- /dev/null +++ b/backend/windmill-api/src/bedrock.rs @@ -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::(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::(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 = 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 { + 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> + + Send + + 'static, + model: String, +) -> impl futures::Stream> + 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, // index -> (id, name, args) + buffer: Vec, // 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::(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")) + ])) +} diff --git a/backend/windmill-api/src/lib.rs b/backend/windmill-api/src/lib.rs index cd2f963680..3a83b11ce8 100644 --- a/backend/windmill-api/src/lib.rs +++ b/backend/windmill-api/src/lib.rs @@ -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; diff --git a/backend/windmill-common/Cargo.toml b/backend/windmill-common/Cargo.toml index 674c48ffd3..123d1a5ea4 100644 --- a/backend/windmill-common/Cargo.toml +++ b/backend/windmill-common/Cargo.toml @@ -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 diff --git a/backend/windmill-common/src/ai_providers.rs b/backend/windmill-common/src/ai_providers.rs index 46d5d687c3..58d2e041e4 100644 --- a/backend/windmill-common/src/ai_providers.rs +++ b/backend/windmill-common/src/ai_providers.rs @@ -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, db: &DB) -> Result { + pub async fn get_base_url( + &self, + resource_base_url: Option, + region: Option, + db: &DB, + ) -> Result { 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()) + )), } } diff --git a/backend/windmill-worker/Cargo.toml b/backend/windmill-worker/Cargo.toml index d7e0f66685..19ebb88db5 100644 --- a/backend/windmill-worker/Cargo.toml +++ b/backend/windmill-worker/Cargo.toml @@ -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 diff --git a/backend/windmill-worker/src/ai/image_handler.rs b/backend/windmill-worker/src/ai/image_handler.rs index c668ded17f..946c922fed 100644 --- a/backend/windmill-worker/src/ai/image_handler.rs +++ b/backend/windmill-worker/src/ai/image_handler.rs @@ -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, 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) +} diff --git a/backend/windmill-worker/src/ai/providers/bedrock.rs b/backend/windmill-worker/src/ai/providers/bedrock.rs new file mode 100644 index 0000000000..18d08d0a41 --- /dev/null +++ b/backend/windmill-worker/src/ai/providers/bedrock.rs @@ -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 { + 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(error: &aws_sdk_bedrockruntime::error::SdkError) -> 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, Vec), 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), Error> { + if !url.starts_with("data:") { + return Err(Error::internal_err("Image URL must be a data URL")); + } + + // Parse data:image/png;base64, + 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, 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 { + 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 { + 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 { + 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::(&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, 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, + max_tokens: Option, +) -> Option { + 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, Vec), 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 { + 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 { + 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 { + 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) -> Vec { + 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, + max_tokens: Option, + api_key: &str, + region: &str, + should_stream: bool, + stream_event_processor: Option, + client: &AuthedClient, + workspace_id: &str, + structured_output_tool_name: Option<&str>, + ) -> Result { + // 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, 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, + system_prompts: Vec, + inference_config: Option, + tool_config: Option, + ) -> Result { + 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, + system_prompts: Vec, + inference_config: Option, + tool_config: Option, + stream_event_processor: Option, + ) -> Result { + // 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 = HashMap::new(); + let mut current_tool_use_id: Option = 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) = ¤t_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) + }, + }) + } +} diff --git a/backend/windmill-worker/src/ai/providers/google_ai.rs b/backend/windmill-worker/src/ai/providers/google_ai.rs index b94d4f29e5..0248c8d7e4 100644 --- a/backend/windmill-worker/src/ai/providers/google_ai.rs +++ b/backend/windmill-worker/src/ai/providers/google_ai.rs @@ -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 => { diff --git a/backend/windmill-worker/src/ai/providers/mod.rs b/backend/windmill-worker/src/ai/providers/mod.rs index 13cf766e28..e86d70650a 100644 --- a/backend/windmill-worker/src/ai/providers/mod.rs +++ b/backend/windmill-worker/src/ai/providers/mod.rs @@ -1,3 +1,4 @@ +pub mod bedrock; pub mod google_ai; pub mod openai; pub mod openrouter; diff --git a/backend/windmill-worker/src/ai/providers/openai.rs b/backend/windmill-worker/src/ai/providers/openai.rs index d764cc11ad..64a6e4a1b8 100644 --- a/backend/windmill-worker/src/ai/providers/openai.rs +++ b/backend/windmill-worker/src/ai/providers/openai.rs @@ -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, 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 { - 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", diff --git a/backend/windmill-worker/src/ai/providers/openrouter.rs b/backend/windmill-worker/src/ai/providers/openrouter.rs index 9e22a63552..f3e215cada 100644 --- a/backend/windmill-worker/src/ai/providers/openrouter.rs +++ b/backend/windmill-worker/src/ai/providers/openrouter.rs @@ -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) } diff --git a/backend/windmill-worker/src/ai/query_builder.rs b/backend/windmill-worker/src/ai/query_builder.rs index e332c21f73..603f428c2e 100644 --- a/backend/windmill-worker/src/ai/query_builder.rs +++ b/backend/windmill-worker/src/ai/query_builder.rs @@ -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( diff --git a/backend/windmill-worker/src/ai/types.rs b/backend/windmill-worker/src/ai/types.rs index 457eb3b215..eb22066436 100644 --- a/backend/windmill-worker/src/ai/types.rs +++ b/backend/windmill-worker/src/ai/types.rs @@ -126,6 +126,7 @@ pub struct ProviderResource { pub api_key: String, #[serde(alias = "baseUrl")] pub base_url: Option, + pub region: Option, } #[derive(Deserialize, Debug)] @@ -146,9 +147,17 @@ impl ProviderWithResource { pub async fn get_base_url(&self, db: &DB) -> Result { 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)] diff --git a/backend/windmill-worker/src/ai/utils.rs b/backend/windmill-worker/src/ai/utils.rs index 947b56f549..34b1ad69c4 100644 --- a/backend/windmill-worker/src/ai/utils.rs +++ b/backend/windmill-worker/src/ai/utils.rs @@ -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 diff --git a/backend/windmill-worker/src/ai_executor.rs b/backend/windmill-worker/src/ai_executor.rs index 481717a207..aff089b82b 100644 --- a/backend/windmill-worker/src/ai_executor.rs +++ b/backend/windmill-worker/src/ai_executor.rs @@ -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 = 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(|_| "".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(|_| "".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); } } } diff --git a/frontend/src/lib/components/copilot/lib.ts b/frontend/src/lib/components/copilot/lib.ts index e589f2d751..4914e538c5 100644 --- a/frontend/src/lib/components/copilot/lib.ts +++ b/frontend/src/lib/components/copilot/lib.ts @@ -76,6 +76,10 @@ export const AI_PROVIDERS: Record = { 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 { - 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