diff --git a/backend/CLAUDE.md b/backend/CLAUDE.md index 78faa0bf77..1fe509a0d6 100644 --- a/backend/CLAUDE.md +++ b/backend/CLAUDE.md @@ -10,16 +10,3 @@ 1. Update database schema with migration if necessary 2. Update backend/windmill-api/openapi.yaml after modifying API endpoints - -## Querying the Database - -To query the database directly, use psql with the following connection string: - -```bash -psql postgres://postgres:changeme@localhost:5432/windmill -``` - -This can be helpful for: -- Inspecting database state during development -- Testing queries before implementing them in Rust -- Debugging data-related issues diff --git a/backend/windmill-api/openapi.yaml b/backend/windmill-api/openapi.yaml index 86b0a840f0..8fa61418ce 100644 --- a/backend/windmill-api/openapi.yaml +++ b/backend/windmill-api/openapi.yaml @@ -572,12 +572,12 @@ paths: use_case: type: string responses: - "200": + '200': description: Onboarding data submitted successfully content: application/json: schema: - type: string + type: string /w/{workspace}/users/delete/{username}: delete: @@ -15046,7 +15046,8 @@ 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 @@ -15054,7 +15055,8 @@ 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 @@ -15349,7 +15351,6 @@ components: groq, openrouter, togetherai, - aws_bedrock, customai, ] @@ -16501,30 +16502,30 @@ components: ScriptLang: type: string enum: [ - python3, - deno, - go, - bash, - powershell, - postgresql, - mysql, - bigquery, - snowflake, - mssql, - oracledb, - graphql, - nativets, - bun, - php, - rust, - ansible, - csharp, - nu, - java, - ruby, - duckdb, - # for related places search: ADD_NEW_LANG - ] + python3, + deno, + go, + bash, + powershell, + postgresql, + mysql, + bigquery, + snowflake, + mssql, + oracledb, + graphql, + nativets, + bun, + php, + rust, + ansible, + csharp, + nu, + java, + ruby, + duckdb, + # for related places search: ADD_NEW_LANG + ] Preview: type: object diff --git a/backend/windmill-api/src/ai.rs b/backend/windmill-api/src/ai.rs index 2bcb6cb699..e4d07cfdc6 100644 --- a/backend/windmill-api/src/ai.rs +++ b/backend/windmill-api/src/ai.rs @@ -1,20 +1,17 @@ use crate::db::{ApiAuthed, DB}; use axum::{body::Bytes, extract::Path, response::IntoResponse, routing::post, Extension, Router}; -use bytes; -use futures; use http::{HeaderMap, Method}; 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 uuid; 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() @@ -183,34 +180,15 @@ 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); - // Handle AWS Bedrock transformation - let (url, body) = if is_bedrock && method != Method::GET { - let (model, transformed_body, is_streaming) = Self::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 url = if is_azure && method != Method::GET { let model = AIProvider::extract_model_from_body(&body)?; - let azure_url = AIProvider::build_azure_openai_url(base_url, &model, path); - (azure_url, body) + AIProvider::build_azure_openai_url(base_url, &model, path) } else if is_anthropic_sdk { let truncated_base_url = base_url.trim_end_matches("/v1"); - let anthropic_url = format!("{}/{}", truncated_base_url, path); - (anthropic_url, body) + format!("{}/{}", truncated_base_url, path) } else { - let default_url = format!("{}/{}", base_url, path); - (default_url, body) + format!("{}/{}", base_url, path) }; tracing::debug!("AI request URL: {}", url); @@ -275,604 +253,6 @@ impl AIRequestConfig { .map_err(|e| Error::internal_err(format!("Failed to reserialize request body: {}", e)))? .into()) } - - /// Transform OpenAI format request to AWS Bedrock Converse format - /// Returns: (model_id, transformed_body, is_streaming) - fn transform_openai_to_bedrock(body: &[u8]) -> Result<(String, Bytes, bool)> { - use serde_json::Value; - - // Parse the OpenAI request - let openai_req: Value = serde_json::from_slice(body) - .map_err(|e| Error::internal_err(format!("Failed to parse OpenAI request: {}", e)))?; - - // Extract model and streaming flag - let model = openai_req["model"] - .as_str() - .ok_or_else(|| Error::BadRequest("Missing 'model' field in request".to_string()))? - .to_string(); - - let is_streaming = openai_req["stream"].as_bool().unwrap_or(false); - - // Build Bedrock request - let mut bedrock_req = serde_json::json!({}); - - // Transform messages - if let Some(messages) = openai_req["messages"].as_array() { - let mut system_messages = Vec::new(); - let mut conversation_messages = Vec::new(); - - for msg in messages { - let role = msg["role"].as_str().unwrap_or(""); - - match role { - "system" => { - // Extract system messages to separate array - if let Some(content) = msg["content"].as_str() { - system_messages.push(serde_json::json!({"text": content})); - } - } - "user" | "assistant" => { - // Normalize content to array format - let mut content = if let Some(text) = msg["content"].as_str() { - // Simple string → array of content blocks - vec![serde_json::json!({"text": text})] - } else if let Some(content_array) = msg["content"].as_array() { - // Already an array - transform each item - content_array - .iter() - .filter_map(|item| { - if let Some(text) = item["text"].as_str() { - Some(serde_json::json!({"text": text})) - } else if item["type"].as_str() == Some("text") { - Some(serde_json::json!({"text": item["text"]})) - } else if item["type"].as_str() == Some("image_url") { - // Transform image_url format if needed - // For now, pass through - may need more sophisticated handling - Some(item.clone()) - } else { - None - } - }) - .collect() - } else { - vec![] - }; - - // Handle tool_calls for assistant messages (OpenAI → Bedrock toolUse) - if role == "assistant" { - if let Some(tool_calls) = msg["tool_calls"].as_array() { - for tool_call in tool_calls { - if tool_call["type"].as_str() == Some("function") { - let tool_use_id = tool_call["id"].as_str().unwrap_or(""); - let function_name = - tool_call["function"]["name"].as_str().unwrap_or(""); - let arguments_str = tool_call["function"]["arguments"] - .as_str() - .unwrap_or("{}"); - - // Parse arguments JSON string to object - let input = serde_json::from_str::(arguments_str) - .unwrap_or(serde_json::json!({})); - - content.push(serde_json::json!({ - "toolUse": { - "toolUseId": tool_use_id, - "name": function_name, - "input": input - } - })); - } - } - } - } - - // Only add message if it has content - if !content.is_empty() { - conversation_messages.push(serde_json::json!({ - "role": role, - "content": content - })); - } - } - "tool" => { - // Transform tool response to Bedrock format - let tool_call_id = msg["tool_call_id"].as_str().unwrap_or(""); - let content = msg["content"].as_str().unwrap_or(""); - - // Try to parse content as JSON - // Bedrock requires json field to be an object, not a primitive or array - let tool_result_content = - if let Ok(json_content) = serde_json::from_str::(content) { - if json_content.is_object() { - vec![serde_json::json!({"json": json_content})] - } else { - // Wrap primitives and arrays in an object - vec![serde_json::json!({"json": {"result": json_content}})] - } - } else { - vec![serde_json::json!({"text": content})] - }; - - conversation_messages.push(serde_json::json!({ - "role": "user", - "content": [{ - "toolResult": { - "toolUseId": tool_call_id, - "content": tool_result_content - } - }] - })); - } - _ => {} - } - } - - if !system_messages.is_empty() { - bedrock_req["system"] = Value::Array(system_messages); - } - bedrock_req["messages"] = Value::Array(conversation_messages); - } - - // Transform inference parameters - let mut inference_config = serde_json::json!({}); - if let Some(max_tokens) = openai_req["max_tokens"].as_i64() { - inference_config["maxTokens"] = Value::Number(max_tokens.into()); - } - if let Some(temperature) = openai_req["temperature"].as_f64() { - inference_config["temperature"] = serde_json::json!(temperature); - } - if let Some(top_p) = openai_req["top_p"].as_f64() { - inference_config["topP"] = serde_json::json!(top_p); - } - if let Some(stop) = openai_req["stop"].as_array() { - let stop_sequences: Vec = stop - .iter() - .filter_map(|s| s.as_str().map(|s| s.to_string())) - .collect(); - if !stop_sequences.is_empty() { - inference_config["stopSequences"] = - Value::Array(stop_sequences.into_iter().map(Value::String).collect()); - } - } - if !inference_config.as_object().unwrap().is_empty() { - bedrock_req["inferenceConfig"] = inference_config; - } - - // Transform tools if present - if let Some(tools) = openai_req["tools"].as_array() { - let mut bedrock_tools = Vec::new(); - - for tool in tools { - if tool["type"].as_str() == Some("function") { - if let Some(function) = tool["function"].as_object() { - bedrock_tools.push(serde_json::json!({ - "toolSpec": { - "name": function.get("name"), - "description": function.get("description") - .and_then(|v| v.as_str()) - .filter(|s| !s.is_empty()) - .unwrap_or("Tool function"), - "inputSchema": { - "json": function.get("parameters") - } - } - })); - } - } - } - - if !bedrock_tools.is_empty() { - let mut tool_config = serde_json::json!({ - "tools": bedrock_tools - }); - - // Transform tool_choice - if let Some(tool_choice) = openai_req.get("tool_choice") { - if tool_choice == "auto" { - tool_config["toolChoice"] = serde_json::json!({"auto": {}}); - } else if tool_choice == "required" { - tool_config["toolChoice"] = serde_json::json!({"any": {}}); - } else if let Some(obj) = tool_choice.as_object() { - if obj.get("type").and_then(|v| v.as_str()) == Some("function") { - if let Some(function) = obj.get("function").and_then(|v| v.as_object()) - { - if let Some(name) = function.get("name").and_then(|v| v.as_str()) { - tool_config["toolChoice"] = serde_json::json!({ - "tool": {"name": name} - }); - } - } - } - } - } - - bedrock_req["toolConfig"] = tool_config; - } - } - - let transformed_body = serde_json::to_vec(&bedrock_req) - .map_err(|e| { - Error::internal_err(format!("Failed to serialize Bedrock request: {}", e)) - })? - .into(); - - Ok((model, transformed_body, is_streaming)) - } - - /// Transform AWS Bedrock Converse response to OpenAI format - async fn transform_bedrock_to_openai( - response: reqwest::Response, - model: String, - ) -> Result { - use serde_json::Value; - - let bedrock_resp: Value = response - .json() - .await - .map_err(|e| Error::internal_err(format!("Failed to parse Bedrock response: {}", e)))?; - - // Generate unique ID and timestamp - let id = format!("chatcmpl-{}", uuid::Uuid::new_v4().simple()); - let created = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs(); - - // Extract stop reason and map to finish_reason - let stop_reason = bedrock_resp["stopReason"].as_str().unwrap_or("end_turn"); - let finish_reason = match stop_reason { - "end_turn" => "stop", - "max_tokens" => "length", - "tool_use" => "tool_calls", - "stop_sequence" => "stop", - "guardrail_intervened" | "content_filtered" => "content_filter", - _ => "stop", - }; - - // Extract message content - let message_content = &bedrock_resp["output"]["message"]["content"]; - let mut text_content = String::new(); - let mut tool_calls = Vec::new(); - - if let Some(content_array) = message_content.as_array() { - for (_index, block) in content_array.iter().enumerate() { - if let Some(text) = block["text"].as_str() { - text_content.push_str(text); - } else if let Some(tool_use) = block.get("toolUse") { - // Transform tool use to OpenAI tool_calls format - let tool_call_id = tool_use["toolUseId"].as_str().unwrap_or(""); - let name = tool_use["name"].as_str().unwrap_or(""); - let input = &tool_use["input"]; - - tool_calls.push(serde_json::json!({ - "id": tool_call_id, - "type": "function", - "function": { - "name": name, - "arguments": serde_json::to_string(input).unwrap_or_default() - } - })); - } - } - } - - // Build the message - let message = if !tool_calls.is_empty() { - serde_json::json!({ - "role": "assistant", - "content": if text_content.is_empty() { Value::Null } else { Value::String(text_content) }, - "tool_calls": tool_calls - }) - } else { - serde_json::json!({ - "role": "assistant", - "content": text_content - }) - }; - - // Extract usage information - let usage = if let Some(usage_data) = bedrock_resp.get("usage") { - serde_json::json!({ - "prompt_tokens": usage_data["inputTokens"].as_i64().unwrap_or(0), - "completion_tokens": usage_data["outputTokens"].as_i64().unwrap_or(0), - "total_tokens": usage_data["totalTokens"].as_i64().unwrap_or(0) - }) - } else { - serde_json::json!({ - "prompt_tokens": 0, - "completion_tokens": 0, - "total_tokens": 0 - }) - }; - - // Build OpenAI-format response - let openai_resp = serde_json::json!({ - "id": id, - "object": "chat.completion", - "created": created, - "model": model, - "choices": [{ - "index": 0, - "message": message, - "finish_reason": finish_reason - }], - "usage": usage - }); - - let response_body = serde_json::to_vec(&openai_resp) - .map_err(|e| { - Error::internal_err(format!("Failed to serialize OpenAI response: {}", e)) - })? - .into(); - - Ok(response_body) - } - - /// Transform AWS Bedrock streaming response to OpenAI SSE format - /// Bedrock uses AWS event stream binary format, not SSE - fn transform_bedrock_stream_to_openai( - stream: impl futures::Stream> - + Send - + 'static, - model: String, - ) -> impl futures::Stream> + Send { - use futures::stream::StreamExt; - use serde_json::Value; - use std::collections::HashMap; - - let id = format!("chatcmpl-{}", uuid::Uuid::new_v4().simple()); - let created = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs(); - - // State to track partial tool calls and binary buffer - struct StreamState { - id: String, - model: String, - created: u64, - tool_calls: HashMap, // index -> (id, name, args) - buffer: Vec, // Binary buffer for AWS event stream - } - - let state = std::sync::Arc::new(tokio::sync::Mutex::new(StreamState { - id: id.clone(), - model: model.clone(), - created, - tool_calls: HashMap::new(), - buffer: Vec::new(), - })); - - stream - .then(move |chunk_result| { - let state = state.clone(); - async move { - match chunk_result { - Ok(chunk) => { - let mut state = state.lock().await; - state.buffer.extend_from_slice(&chunk); - - let mut events = Vec::new(); - - // Parse AWS event stream messages from buffer - loop { - // Need at least 12 bytes for prelude (8) + prelude CRC (4) - if state.buffer.len() < 12 { - break; - } - - // Read prelude: total_length (4 bytes) + headers_length (4 bytes) - let total_length = u32::from_be_bytes([ - state.buffer[0], - state.buffer[1], - state.buffer[2], - state.buffer[3], - ]) as usize; - - // Check if we have the complete message - if state.buffer.len() < total_length { - break; - } - - let headers_length = u32::from_be_bytes([ - state.buffer[4], - state.buffer[5], - state.buffer[6], - state.buffer[7], - ]) as usize; - - // Skip prelude CRC (4 bytes after prelude) - let headers_start = 12; - let payload_start = headers_start + headers_length; - let payload_end = total_length - 4; // Exclude message CRC - - // Parse headers to extract event type - let mut event_type = None; - let mut pos = headers_start; - while pos < payload_start { - if pos + 1 > state.buffer.len() { - break; - } - let name_len = state.buffer[pos] as usize; - pos += 1; - - if pos + name_len > state.buffer.len() { - break; - } - let name = String::from_utf8_lossy(&state.buffer[pos..pos + name_len]).to_string(); - pos += name_len; - - if pos + 3 > state.buffer.len() { - break; - } - let value_type = state.buffer[pos]; - pos += 1; - let value_len = u16::from_be_bytes([state.buffer[pos], state.buffer[pos + 1]]) as usize; - pos += 2; - - if pos + value_len > state.buffer.len() { - break; - } - - if value_type == 7 && name == ":event-type" { - event_type = Some(String::from_utf8_lossy(&state.buffer[pos..pos + value_len]).to_string()); - } - pos += value_len; - } - - // Extract JSON payload (copy to avoid borrow issues) - let payload = state.buffer[payload_start..payload_end].to_vec(); - - // Remove processed message from buffer - state.buffer.drain(0..total_length); - - // Process the event - if let Some(evt_type) = event_type { - if let Ok(payload_str) = std::str::from_utf8(&payload) { - if let Ok(parsed_data) = serde_json::from_str::(payload_str) { - // Transform based on event type - match evt_type.as_str() { - "messageStart" => { - // No output for messageStart - } - "contentBlockStart" => { - let index = parsed_data["contentBlockIndex"].as_u64().unwrap_or(0) as usize; - - if let Some(tool_use) = parsed_data["start"].get("toolUse") { - let tool_id = tool_use["toolUseId"].as_str().unwrap_or("").to_string(); - let name = tool_use["name"].as_str().unwrap_or("").to_string(); - - state.tool_calls.insert(index, (tool_id.clone(), name.clone(), String::new())); - - // Send initial tool call chunk - let chunk = serde_json::json!({ - "id": state.id, - "object": "chat.completion.chunk", - "created": state.created, - "model": state.model, - "choices": [{ - "index": 0, - "delta": { - "tool_calls": [{ - "index": index, - "id": tool_id, - "type": "function", - "function": { - "name": name, - "arguments": "" - } - }] - }, - "finish_reason": Value::Null - }] - }); - - events.push(Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk)))); - } - } - "contentBlockDelta" => { - let index = parsed_data["contentBlockIndex"].as_u64().unwrap_or(0) as usize; - - if let Some(text) = parsed_data["delta"]["text"].as_str() { - // Text content delta - let chunk = serde_json::json!({ - "id": state.id, - "object": "chat.completion.chunk", - "created": state.created, - "model": state.model, - "choices": [{ - "index": 0, - "delta": { - "content": text - }, - "finish_reason": Value::Null - }] - }); - - events.push(Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk)))); - } else if let Some(tool_use_input) = parsed_data["delta"]["toolUse"]["input"].as_str() { - // Tool use arguments delta - if let Some((_tool_id, _name, ref mut args)) = state.tool_calls.get_mut(&index) { - args.push_str(tool_use_input); - - let chunk = serde_json::json!({ - "id": state.id, - "object": "chat.completion.chunk", - "created": state.created, - "model": state.model, - "choices": [{ - "index": 0, - "delta": { - "tool_calls": [{ - "index": index, - "function": { - "arguments": tool_use_input - } - }] - }, - "finish_reason": Value::Null - }] - }); - - events.push(Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk)))); - } - } - } - "contentBlockStop" => { - // No output needed - } - "messageStop" => { - let stop_reason = parsed_data["stopReason"].as_str().unwrap_or("end_turn"); - let finish_reason = match stop_reason { - "end_turn" => "stop", - "max_tokens" => "length", - "tool_use" => "tool_calls", - "stop_sequence" => "stop", - "guardrail_intervened" | "content_filtered" => "content_filter", - _ => "stop", - }; - - let chunk = serde_json::json!({ - "id": state.id, - "object": "chat.completion.chunk", - "created": state.created, - "model": state.model, - "choices": [{ - "index": 0, - "delta": {}, - "finish_reason": finish_reason - }] - }); - - events.push(Ok(bytes::Bytes::from(format!("data: {}\n\n", chunk)))); - } - "metadata" => { - // Could include usage info here if needed - } - _ => {} - } - } - } - } - } // end loop - - events - } - Err(e) => { - vec![Err(std::io::Error::new( - std::io::ErrorKind::Other, - e.to_string(), - ))] - } - } - } - }) - .flat_map(|events| futures::stream::iter(events)) - .chain(futures::stream::iter(vec![ - // Send [DONE] at the end - Ok(bytes::Bytes::from("data: [DONE]\n\n")) - ])) - } } #[derive(Clone, Debug)] @@ -1066,18 +446,6 @@ async fn proxy( } }; - // Extract model and streaming flag for Bedrock transformation (only for POST requests) - let (model_for_transform, is_streaming) = - if matches!(provider, AIProvider::AWSBedrock) && method == Method::POST { - let parsed: serde_json::Value = serde_json::from_slice(&body) - .map_err(|e| Error::internal_err(format!("Failed to parse request body: {}", e)))?; - let model = parsed["model"].as_str().unwrap_or("").to_string(); - let is_streaming = parsed["stream"].as_bool().unwrap_or(false); - (Some(model), is_streaming) - } else { - (None, false) - }; - let request = request_config.prepare_request(&provider, &ai_path, method, headers, body)?; let response = request.send().await.map_err(to_anyhow)?; @@ -1101,47 +469,8 @@ async fn proxy( return Err(Error::AIError(err_msg)); } - // Transform Bedrock responses back to OpenAI format - if matches!(provider, AIProvider::AWSBedrock) && model_for_transform.is_some() { - let model = model_for_transform.unwrap(); - - 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 = - AIRequestConfig::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 = - AIRequestConfig::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))) - } + let status_code = response.status(); + let headers = response.headers().clone(); + let stream = response.bytes_stream(); + Ok((status_code, headers, axum::body::Body::from_stream(stream))) } diff --git a/backend/windmill-common/src/ai_providers.rs b/backend/windmill-common/src/ai_providers.rs index dee7aa2c4b..46d5d687c3 100644 --- a/backend/windmill-common/src/ai_providers.rs +++ b/backend/windmill-common/src/ai_providers.rs @@ -27,8 +27,6 @@ pub enum AIProvider { OpenRouter, TogetherAI, CustomAI, - #[serde(rename = "aws_bedrock")] - AWSBedrock, } impl AIProvider { @@ -66,7 +64,7 @@ impl AIProvider { AIProvider::TogetherAI => Ok("https://api.together.xyz/v1".to_string()), AIProvider::Anthropic => Ok("https://api.anthropic.com/v1".to_string()), AIProvider::Mistral => Ok("https://api.mistral.ai/v1".to_string()), - p @ (AIProvider::CustomAI | AIProvider::AzureOpenAI | AIProvider::AWSBedrock) => { + p @ (AIProvider::CustomAI | AIProvider::AzureOpenAI) => { if let Some(base_url) = resource_base_url { Ok(base_url) } else { diff --git a/backend/windmill-worker/src/ai/image_handler.rs b/backend/windmill-worker/src/ai/image_handler.rs index 946c922fed..c668ded17f 100644 --- a/backend/windmill-worker/src/ai/image_handler.rs +++ b/backend/windmill-worker/src/ai/image_handler.rs @@ -4,8 +4,6 @@ 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, @@ -68,53 +66,3 @@ pub async fn download_and_encode_s3_image( Ok((mime_type.to_string(), base64_data)) } - -/// Prepare messages for API by converting S3Objects to base64 ImageUrls -pub async fn prepare_messages_for_api( - messages: &[OpenAIMessage], - client: &AuthedClient, - workspace_id: &str, -) -> Result, Error> { - let mut prepared_messages = Vec::new(); - - for message in messages { - let mut prepared_message = message.clone(); - - if let Some(content) = &message.content { - match content { - OpenAIContent::Text(text) => { - prepared_message.content = Some(OpenAIContent::Text(text.clone())); - } - OpenAIContent::Parts(parts) => { - let mut prepared_content = Vec::new(); - - for part in parts { - match part { - ContentPart::S3Object { s3_object } => { - // Convert S3Object to base64 image URL - let (mime_type, image_bytes) = - download_and_encode_s3_image(s3_object, client, workspace_id) - .await?; - prepared_content.push(ContentPart::ImageUrl { - image_url: ImageUrlData { - url: format!("data:{};base64,{}", mime_type, image_bytes), - }, - }); - } - other => { - // Keep Text and ImageUrl as-is - prepared_content.push(other.clone()); - } - } - } - - prepared_message.content = Some(OpenAIContent::Parts(prepared_content)); - } - } - } - - prepared_messages.push(prepared_message); - } - - Ok(prepared_messages) -} diff --git a/backend/windmill-worker/src/ai/providers/bedrock.rs b/backend/windmill-worker/src/ai/providers/bedrock.rs deleted file mode 100644 index 3e95d1aaac..0000000000 --- a/backend/windmill-worker/src/ai/providers/bedrock.rs +++ /dev/null @@ -1,588 +0,0 @@ -use async_trait::async_trait; -use futures::StreamExt; -use serde::Deserialize; -use serde_json::Value; -use windmill_common::{client::AuthedClient, error::Error}; - -use crate::ai::{ - image_handler::prepare_messages_for_api, - providers::openai::{OpenAIFunction, OpenAIToolCall}, - query_builder::{BuildRequestArgs, ParsedResponse, QueryBuilder, StreamEventProcessor}, - types::*, -}; - -// Bedrock-specific response types -#[derive(Deserialize)] -struct BedrockResponse { - output: BedrockOutput, - #[allow(unused)] - #[serde(rename = "stopReason")] - stop_reason: Option, -} - -#[derive(Deserialize)] -struct BedrockOutput { - message: BedrockMessage, -} - -#[derive(Deserialize)] -struct BedrockMessage { - #[allow(unused)] - role: String, - content: Vec, -} - -#[derive(Deserialize)] -#[serde(untagged)] -enum BedrockContent { - Text { - text: String, - }, - ToolUse { - #[serde(rename = "toolUse")] - tool_use: ToolUse, - }, -} - -#[derive(Deserialize)] -struct ToolUse { - #[serde(rename = "toolUseId")] - tool_use_id: String, - name: String, - input: Value, -} - -pub struct BedrockQueryBuilder; - -impl BedrockQueryBuilder { - pub fn new() -> Self { - Self - } - - /// Transform OpenAI format messages to Bedrock Converse format - fn transform_messages_to_bedrock( - messages: &[OpenAIMessage], - ) -> Result<(Vec, Vec), Error> { - let mut system_messages = Vec::new(); - let mut conversation_messages = Vec::new(); - - for msg in messages { - let role = &msg.role; - - match role.as_str() { - "system" => { - // Extract system messages - if let Some(content) = &msg.content { - let text = match content { - OpenAIContent::Text(t) => t.clone(), - OpenAIContent::Parts(parts) => { - // Extract text from parts - parts - .iter() - .filter_map(|part| match part { - ContentPart::Text { text } => Some(text.clone()), - _ => None, - }) - .collect::>() - .join(" ") - } - }; - system_messages.push(serde_json::json!({"text": text})); - } - } - "user" | "assistant" => { - let mut content = Vec::new(); - - // Handle message content - if let Some(msg_content) = &msg.content { - match msg_content { - OpenAIContent::Text(text) => { - content.push(serde_json::json!({"text": text})); - } - OpenAIContent::Parts(parts) => { - for part in parts { - match part { - ContentPart::Text { text } => { - content.push(serde_json::json!({"text": text})); - } - ContentPart::ImageUrl { image_url } => { - // Bedrock image format - extract base64 from data URL - let url = &image_url.url; - if url.starts_with("data:") { - // Parse data:image/png;base64, - if let Some(base64_start) = url.find("base64,") { - let base64_data = &url[base64_start + 7..]; - let mime_type = url - .split(';') - .next() - .and_then(|s| s.strip_prefix("data:")) - .unwrap_or("image/png"); - - content.push(serde_json::json!({ - "image": { - "format": mime_type.split('/').last().unwrap_or("png"), - "source": { - "bytes": base64_data - } - } - })); - } - } - } - _ => {} - } - } - } - } - } - - // Handle tool_calls for assistant messages - if role == "assistant" { - if let Some(tool_calls) = &msg.tool_calls { - for tool_call in tool_calls { - if tool_call.r#type == "function" { - // Parse arguments JSON string to object - let input = serde_json::from_str::( - &tool_call.function.arguments, - ) - .unwrap_or(serde_json::json!({})); - - content.push(serde_json::json!({ - "toolUse": { - "toolUseId": tool_call.id, - "name": tool_call.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_ref().map(|s| s.as_str()).unwrap_or(""); - let content_text = match &msg.content { - Some(OpenAIContent::Text(t)) => t.clone(), - Some(OpenAIContent::Parts(parts)) => parts - .iter() - .filter_map(|part| match part { - ContentPart::Text { text } => Some(text.clone()), - _ => None, - }) - .collect::>() - .join(" "), - None => String::new(), - }; - - // Parse content as JSON if possible, otherwise use as text - // Bedrock requires json field to be an object, not a primitive or array - let tool_result_content = - if let Ok(parsed) = serde_json::from_str::(&content_text) { - if parsed.is_object() { - parsed - } else { - // Wrap primitives and arrays in an object - serde_json::json!({"result": parsed}) - } - } else { - serde_json::json!({"result": content_text}) - }; - - // Bedrock requires toolResult to be in a user message - conversation_messages.push(serde_json::json!({ - "role": "user", - "content": [{ - "toolResult": { - "toolUseId": tool_call_id, - "content": [{"json": tool_result_content}] - } - }] - })); - } - _ => { - // Skip unknown roles - } - } - } - - Ok((system_messages, conversation_messages)) - } - - /// Transform Bedrock response to OpenAI format - fn transform_bedrock_response_to_openai( - bedrock_response: BedrockResponse, - ) -> Result { - let mut content_text = String::new(); - let mut tool_calls = Vec::new(); - - for content_item in bedrock_response.output.message.content { - match content_item { - BedrockContent::Text { text } => { - if !content_text.is_empty() { - content_text.push(' '); - } - content_text.push_str(&text); - } - BedrockContent::ToolUse { tool_use } => { - // Convert Bedrock toolUse to OpenAI tool_call - let arguments = - serde_json::to_string(&tool_use.input).unwrap_or_else(|_| "{}".to_string()); - - tool_calls.push(OpenAIToolCall { - id: tool_use.tool_use_id, - function: OpenAIFunction { name: tool_use.name, arguments }, - r#type: "function".to_string(), - }); - } - } - } - - Ok(ParsedResponse::Text { - content: if content_text.is_empty() { - None - } else { - Some(content_text) - }, - tool_calls, - events_str: None, - }) - } -} - -#[async_trait] -impl QueryBuilder for BedrockQueryBuilder { - fn supports_tools_with_output_type(&self, output_type: &OutputType) -> bool { - // Bedrock supports tools for text output - matches!(output_type, OutputType::Text) - } - - fn supports_streaming(&self) -> bool { - // Bedrock supports streaming - true - } - - async fn build_request( - &self, - args: &BuildRequestArgs<'_>, - client: &AuthedClient, - workspace_id: &str, - _stream: bool, - ) -> Result { - // Only support text output for now - if !matches!(args.output_type, OutputType::Text) { - return Err(Error::internal_err( - "Bedrock only supports text output type".to_string(), - )); - } - - // Prepare messages first (converts S3Objects to ImageUrls) - let prepared_messages = prepare_messages_for_api(args.messages, client, workspace_id).await?; - - // Transform messages - let (system_messages, conversation_messages) = - Self::transform_messages_to_bedrock(&prepared_messages)?; - - // Build Bedrock request - let mut bedrock_req = serde_json::json!({ - "messages": conversation_messages, - }); - - // Add system messages if any - if !system_messages.is_empty() { - bedrock_req["system"] = serde_json::json!(system_messages); - } - - // Add inference configuration - let mut inference_config = serde_json::json!({}); - if let Some(temp) = args.temperature { - inference_config["temperature"] = serde_json::json!(temp); - } - if let Some(max_tokens) = args.max_tokens { - inference_config["maxTokens"] = serde_json::json!(max_tokens); - } - - if !inference_config.as_object().unwrap().is_empty() { - bedrock_req["inferenceConfig"] = inference_config; - } - - // Add tools if provided - if let Some(tools) = args.tools { - let bedrock_tools: Vec = tools - .iter() - .map(|tool| { - // Parse the parameters from RawValue - let params: Value = serde_json::from_str(tool.function.parameters.get()) - .unwrap_or(serde_json::json!({})); - - serde_json::json!({ - "toolSpec": { - "name": tool.function.name, - "description": tool.function.description.as_ref().map(|s| s.as_str()).unwrap_or("Tool function"), - "inputSchema": { - "json": params - } - } - }) - }) - .collect(); - - bedrock_req["toolConfig"] = serde_json::json!({ - "tools": bedrock_tools, - }); - - // Handle structured output schema - let has_output_properties = args - .output_schema - .as_ref() - .and_then(|schema| schema.properties.as_ref()) - .map(|props| !props.is_empty()) - .unwrap_or(false); - - if has_output_properties { - bedrock_req["toolConfig"]["toolChoice"] = serde_json::json!({ - "any": {} - }); - } - } - - serde_json::to_string(&bedrock_req) - .map_err(|e| Error::internal_err(format!("Failed to serialize Bedrock request: {}", e))) - } - - async fn parse_response(&self, response: reqwest::Response) -> Result { - let response_text = response - .text() - .await - .map_err(|e| Error::internal_err(format!("Failed to read response text: {}", e)))?; - - let bedrock_response: BedrockResponse = - serde_json::from_str(&response_text).map_err(|e| { - Error::internal_err(format!( - "Failed to parse Bedrock response: {}. Raw response: {}", - e, response_text - )) - })?; - - Self::transform_bedrock_response_to_openai(bedrock_response) - } - - async fn parse_streaming_response( - &self, - response: reqwest::Response, - stream_event_processor: StreamEventProcessor, - ) -> Result { - let mut stream = response.bytes_stream(); - let mut buffer: Vec = Vec::new(); - let mut accumulated_content = String::new(); - let mut accumulated_tool_calls: std::collections::HashMap = - std::collections::HashMap::new(); - let mut events_str = String::new(); - - while let Some(chunk) = stream.next().await { - let chunk = chunk.map_err(|e| Error::internal_err(format!("Stream error: {}", e)))?; - buffer.extend_from_slice(&chunk); - - // Parse AWS event stream binary format - while buffer.len() >= 12 { - // Need at least prelude + CRC - // Read prelude - let total_length = - u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]) as usize; - - // Check if we have the complete message - if buffer.len() < total_length { - break; - } - - let headers_length = - u32::from_be_bytes([buffer[4], buffer[5], buffer[6], buffer[7]]) as usize; - - let headers_start = 12; // After prelude (8 bytes) + prelude CRC (4 bytes) - let payload_start = headers_start + headers_length; - let payload_end = total_length - 4; // Before message CRC (4 bytes) - - // Parse headers to extract event type - let mut event_type = None; - let mut pos = headers_start; - while pos < payload_start && pos < buffer.len() { - if pos + 1 > buffer.len() { - break; - } - let name_len = buffer[pos] as usize; - pos += 1; - - if pos + name_len > buffer.len() { - break; - } - let name = String::from_utf8_lossy(&buffer[pos..pos + name_len]).to_string(); - pos += name_len; - - if pos + 3 > buffer.len() { - break; - } - let value_type = buffer[pos]; - pos += 1; - let value_len = u16::from_be_bytes([buffer[pos], buffer[pos + 1]]) as usize; - pos += 2; - - if pos + value_len > buffer.len() { - break; - } - - if value_type == 7 && name == ":event-type" { - event_type = Some( - String::from_utf8_lossy(&buffer[pos..pos + value_len]).to_string(), - ); - } - pos += value_len; - } - - // Extract and parse JSON payload - if payload_start < payload_end && payload_end <= buffer.len() { - let payload = &buffer[payload_start..payload_end]; - - if let Ok(event_data) = serde_json::from_slice::(payload) { - // Handle different event types - match event_type.as_deref() { - Some("contentBlockStart") => { - // Tool use started - if let Some(tool_use) = - event_data.get("start").and_then(|s| s.get("toolUse")) - { - if let Some(tool_use_id) = - tool_use.get("toolUseId").and_then(|id| id.as_str()) - { - let name = tool_use - .get("name") - .and_then(|n| n.as_str()) - .unwrap_or("") - .to_string(); - - accumulated_tool_calls.insert( - tool_use_id.to_string(), - OpenAIToolCall { - id: tool_use_id.to_string(), - function: OpenAIFunction { - name, - arguments: String::new(), - }, - r#type: "function".to_string(), - }, - ); - } - } - } - Some("contentBlockDelta") => { - if let Some(delta) = event_data.get("delta") { - // Text delta - if let Some(text) = delta.get("text").and_then(|t| t.as_str()) { - accumulated_content.push_str(text); - - let event = StreamingEvent::TokenDelta { - content: text.to_string(), - }; - stream_event_processor.send(event, &mut events_str).await?; - } - - // Tool use delta (input accumulation) - if let Some(tool_use) = delta.get("toolUse") { - if let Some(input_str) = - tool_use.get("input").and_then(|i| i.as_str()) - { - // Find the tool call being updated (last one added) - if let Some(last_tool_call) = - accumulated_tool_calls.values_mut().last() - { - last_tool_call - .function - .arguments - .push_str(input_str); - } - } - } - } - } - Some("contentBlockStop") => { - // Block completed - nothing to do - } - Some("messageStop") => { - // Message completed - break; - } - Some("metadata") => { - // Usage information - ignore for now - } - _ => { - // Unknown event type - ignore - } - } - } - } - - // Remove processed message from buffer - buffer.drain(0..total_length); - } - } - - // Send tool call events - for tool_call in accumulated_tool_calls.values() { - let event = StreamingEvent::ToolCallArguments { - call_id: tool_call.id.clone(), - function_name: tool_call.function.name.clone(), - arguments: tool_call.function.arguments.clone(), - }; - stream_event_processor.send(event, &mut events_str).await?; - } - - Ok(ParsedResponse::Text { - content: if accumulated_content.is_empty() { - None - } else { - Some(accumulated_content) - }, - tool_calls: accumulated_tool_calls.into_values().collect(), - events_str: Some(events_str), - }) - } - - fn get_endpoint( - &self, - base_url: &str, - model: &str, - output_type: &OutputType, - stream: bool, - ) -> String { - // Bedrock uses different URL structure: /model/{model-id}/converse[-stream] - if !matches!(output_type, OutputType::Text) { - // Image generation not supported yet - return format!("{}/model/{}/converse", base_url, model); - } - - // Use -stream suffix for streaming requests - let endpoint = if stream { - "converse-stream" - } else { - "converse" - }; - format!("{}/model/{}/{}", base_url, model, endpoint) - } - - fn get_auth_headers( - &self, - api_key: &str, - _base_url: &str, - _output_type: &OutputType, - ) -> Vec<(&'static str, String)> { - // Bedrock uses Bearer token authentication - vec![("Authorization", format!("Bearer {}", api_key))] - } -} diff --git a/backend/windmill-worker/src/ai/providers/google_ai.rs b/backend/windmill-worker/src/ai/providers/google_ai.rs index 0248c8d7e4..b94d4f29e5 100644 --- a/backend/windmill-worker/src/ai/providers/google_ai.rs +++ b/backend/windmill-worker/src/ai/providers/google_ai.rs @@ -255,13 +255,7 @@ impl QueryBuilder for GoogleAIQueryBuilder { .await } - fn get_endpoint( - &self, - base_url: &str, - model: &str, - output_type: &OutputType, - _stream: bool, - ) -> String { + fn get_endpoint(&self, base_url: &str, model: &str, output_type: &OutputType) -> String { match output_type { OutputType::Text => format!("{}/chat/completions", base_url), // Use OpenAI-compatible endpoint OutputType::Image => { diff --git a/backend/windmill-worker/src/ai/providers/mod.rs b/backend/windmill-worker/src/ai/providers/mod.rs index e86d70650a..13cf766e28 100644 --- a/backend/windmill-worker/src/ai/providers/mod.rs +++ b/backend/windmill-worker/src/ai/providers/mod.rs @@ -1,4 +1,3 @@ -pub mod bedrock; pub mod google_ai; pub mod openai; pub mod openrouter; diff --git a/backend/windmill-worker/src/ai/providers/openai.rs b/backend/windmill-worker/src/ai/providers/openai.rs index 0d1651c160..d764cc11ad 100644 --- a/backend/windmill-worker/src/ai/providers/openai.rs +++ b/backend/windmill-worker/src/ai/providers/openai.rs @@ -4,11 +4,11 @@ use serde_json; use windmill_common::{ai_providers::AIProvider, client::AuthedClient, error::Error}; use crate::ai::{ - image_handler::{download_and_encode_s3_image, prepare_messages_for_api}, + image_handler::download_and_encode_s3_image, query_builder::{BuildRequestArgs, ParsedResponse, QueryBuilder, StreamEventProcessor}, sse::{OpenAISSEParser, SSEParser}, types::*, - utils::should_use_structured_output_tool, + utils::is_claude_model, }; // OpenAI-specific types @@ -114,6 +114,62 @@ impl OpenAIQueryBuilder { Self { provider_kind } } + pub async fn prepare_messages_for_api( + &self, + messages: &[OpenAIMessage], + client: &AuthedClient, + workspace_id: &str, + ) -> Result, Error> { + let mut prepared_messages = Vec::new(); + + for message in messages { + let mut prepared_message = message.clone(); + + if let Some(content) = &message.content { + match content { + OpenAIContent::Text(text) => { + prepared_message.content = Some(OpenAIContent::Text(text.clone())); + } + OpenAIContent::Parts(parts) => { + let mut prepared_content = Vec::new(); + + for part in parts { + match part { + ContentPart::S3Object { s3_object } => { + // Convert S3Object to base64 image URL + let (mime_type, image_bytes) = download_and_encode_s3_image( + s3_object, + client, + workspace_id, + ) + .await?; + prepared_content.push(ContentPart::ImageUrl { + image_url: ImageUrlData { + url: format!( + "data:{};base64,{}", + mime_type, image_bytes + ), + }, + }); + } + other => { + // Keep Text and ImageUrl as-is + prepared_content.push(other.clone()); + } + } + } + + prepared_message.content = Some(OpenAIContent::Parts(prepared_content)); + } + } + } + + prepared_messages.push(prepared_message); + } + + Ok(prepared_messages) + } + async fn build_text_request( &self, args: &BuildRequestArgs<'_>, @@ -121,8 +177,9 @@ impl OpenAIQueryBuilder { workspace_id: &str, stream: bool, ) -> Result { - let prepared_messages = - prepare_messages_for_api(args.messages, client, workspace_id).await?; + let prepared_messages = self + .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 @@ -146,10 +203,9 @@ impl OpenAIQueryBuilder { None }; - let should_use_structured_output_tool = - should_use_structured_output_tool(&self.provider_kind, args.model); + let is_claude_model = is_claude_model(&args.model); // Force usage of structured output tool for Claude models when structured output provided - let tool_choice = if should_use_structured_output_tool && response_format.is_some() { + let tool_choice = if is_claude_model && response_format.is_some() { Some(ToolChoice::Required) } else { None @@ -349,13 +405,7 @@ impl QueryBuilder for OpenAIQueryBuilder { }) } - fn get_endpoint( - &self, - base_url: &str, - model: &str, - output_type: &OutputType, - _stream: bool, - ) -> String { + fn get_endpoint(&self, base_url: &str, model: &str, output_type: &OutputType) -> String { let path = match output_type { OutputType::Text => "chat/completions", OutputType::Image => "responses", diff --git a/backend/windmill-worker/src/ai/providers/openrouter.rs b/backend/windmill-worker/src/ai/providers/openrouter.rs index f3e215cada..9e22a63552 100644 --- a/backend/windmill-worker/src/ai/providers/openrouter.rs +++ b/backend/windmill-worker/src/ai/providers/openrouter.rs @@ -4,7 +4,6 @@ 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::*, @@ -92,9 +91,11 @@ impl QueryBuilder for OpenRouterQueryBuilder { } OutputType::Image => { // For image generation, we need to add modalities field - // First, prepare the messages - let prepared_messages = - prepare_messages_for_api(args.messages, client, workspace_id).await?; + // 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?; // Check if we need to add response_format for structured output let has_output_properties = args @@ -203,13 +204,7 @@ impl QueryBuilder for OpenRouterQueryBuilder { .await } - fn get_endpoint( - &self, - base_url: &str, - _model: &str, - _output_type: &OutputType, - _stream: bool, - ) -> String { + fn get_endpoint(&self, base_url: &str, _model: &str, _output_type: &OutputType) -> String { // OpenRouter uses the same endpoint for both text and image generation format!("{}/chat/completions", base_url) } diff --git a/backend/windmill-worker/src/ai/query_builder.rs b/backend/windmill-worker/src/ai/query_builder.rs index 88c4faaa87..e332c21f73 100644 --- a/backend/windmill-worker/src/ai/query_builder.rs +++ b/backend/windmill-worker/src/ai/query_builder.rs @@ -7,7 +7,7 @@ use windmill_queue::MiniPulledJob; use crate::{ ai::{ providers::{ - bedrock::BedrockQueryBuilder, google_ai::GoogleAIQueryBuilder, + google_ai::GoogleAIQueryBuilder, openai::{OpenAIQueryBuilder, OpenAIToolCall}, openrouter::OpenRouterQueryBuilder, }, @@ -69,13 +69,7 @@ pub trait QueryBuilder: Send + Sync { } /// Get the API endpoint for this provider - fn get_endpoint( - &self, - base_url: &str, - model: &str, - output_type: &OutputType, - stream: bool, - ) -> String; + fn get_endpoint(&self, base_url: &str, model: &str, output_type: &OutputType) -> String; /// Get the authentication headers for this provider fn get_auth_headers( @@ -93,7 +87,6 @@ pub fn create_query_builder(provider: &ProviderWithResource) -> Box Box::new(GoogleAIQueryBuilder::new()), AIProvider::OpenRouter => Box::new(OpenRouterQueryBuilder::new()), - AIProvider::AWSBedrock => Box::new(BedrockQueryBuilder::new()), _ => Box::new(OpenAIQueryBuilder::new(provider.kind.clone())), // Pass provider kind for Azure handling } } diff --git a/backend/windmill-worker/src/ai/utils.rs b/backend/windmill-worker/src/ai/utils.rs index 34b1ad69c4..947b56f549 100644 --- a/backend/windmill-worker/src/ai/utils.rs +++ b/backend/windmill-worker/src/ai/utils.rs @@ -8,7 +8,6 @@ use std::{ }; use uuid::Uuid; use windmill_common::{ - ai_providers::AIProvider, db::DB, error::Error, flow_conversations::{add_message_to_conversation_tx, MessageType}, @@ -312,9 +311,9 @@ pub fn get_step_name_from_flow( ) } -/// 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 +/// 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") } /// Cleanup MCP clients by gracefully shutting down connections diff --git a/backend/windmill-worker/src/ai_executor.rs b/backend/windmill-worker/src/ai_executor.rs index ed71ccc774..481717a207 100644 --- a/backend/windmill-worker/src/ai_executor.rs +++ b/backend/windmill-worker/src/ai_executor.rs @@ -2,9 +2,9 @@ use crate::ai::tools::{execute_tool_calls, ToolExecutionContext}; use crate::ai::utils::{ add_message_to_conversation, any_tool_needs_previous_result, cleanup_mcp_clients, filter_schema_by_input_transforms, find_unique_tool_name, get_flow_context, - get_flow_job_runnable_and_raw_flow, get_step_name_from_flow, 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, + 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, }; use crate::memory_oss::{read_from_memory, write_to_memory}; use crate::worker_flow::{get_previous_job_result, get_transform_context}; @@ -498,15 +498,14 @@ pub async fn run_agent( .map(|props| !props.is_empty()) .unwrap_or(false); - let should_use_structured_output_tool = - should_use_structured_output_tool(&args.provider.kind, &args.provider.model); + let is_claude_model = is_claude_model(&args.provider.model); let mut used_structured_output_tool = false; let mut structured_output_tool_name: Option = None; // For text output with schema, handle structured output if has_output_properties && output_type == &OutputType::Text { let schema = args.output_schema.as_ref().unwrap(); - if should_use_structured_output_tool { + if is_claude_model { // 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()); @@ -570,12 +569,8 @@ pub async fn run_agent( .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, - should_stream, - ); + 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 timeout = resolve_job_timeout(conn, &job.workspace_id, job.id, job.timeout) diff --git a/frontend/src/lib/components/copilot/lib.ts b/frontend/src/lib/components/copilot/lib.ts index f21ee83756..e589f2d751 100644 --- a/frontend/src/lib/components/copilot/lib.ts +++ b/frontend/src/lib/components/copilot/lib.ts @@ -76,10 +76,6 @@ export const AI_PROVIDERS: Record = { label: 'Together AI', defaultModels: ['meta-llama/Llama-3.3-70B-Instruct-Turbo'] }, - aws_bedrock: { - label: 'AWS Bedrock', - defaultModels: ['amazon.titan-embed-image-v1:0'] - }, customai: { label: 'Custom AI', defaultModels: [] @@ -104,115 +100,18 @@ export async function fetchAvailableModels( provider: AIProvider, signal?: AbortSignal ): Promise { - // Handle AWS Bedrock separately (needs both foundation-models and inference-profiles) - if (provider === 'aws_bedrock') { - const headers = { + const models = await fetch(`${location.origin}${OpenAPI.BASE}/w/${workspace}/ai/proxy/models`, { + signal, + headers: { 'X-Resource-Path': resourcePath, - 'X-Provider': provider + 'X-Provider': provider, + ...(provider === 'anthropic' ? { 'anthropic-version': '2023-06-01' } : {}) } - - // 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') - ) - - // Map models to their invocable IDs - const modelIds = textModels.map((model) => { - const supportsOnDemand = model.inferenceTypesSupported?.includes('ON_DEMAND') - - // If model supports ON_DEMAND, use the model ID directly - if (supportsOnDemand) { - return model.modelId - } - - // Otherwise, find matching inference profile - const matchingProfile = inferenceProfiles.find((profile) => - profile.models.some((m) => m.modelArn === model.modelArn) - ) - - if (matchingProfile) { - return matchingProfile.inferenceProfileId - } - - // Fallback to model ID if no matching profile found (may fail at runtime) - console.warn(`No inference profile found for ${model.modelId}, using direct ID`) - return model.modelId - }) - - // 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) => { @@ -372,8 +271,7 @@ export const PROVIDER_COMPLETION_CONFIG_MAP: Record