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