Compare commits

...

1 Commits

Author SHA1 Message Date
Alex Petric
d3ef73144b experiments with bedrock 2025-04-10 11:08:20 -04:00
9 changed files with 523 additions and 24 deletions

55
backend/Cargo.lock generated
View File

@@ -802,6 +802,7 @@ dependencies = [
"aws-credential-types",
"aws-sigv4",
"aws-smithy-async",
"aws-smithy-eventstream",
"aws-smithy-http",
"aws-smithy-runtime",
"aws-smithy-runtime-api",
@@ -818,6 +819,32 @@ dependencies = [
"uuid",
]
[[package]]
name = "aws-sdk-bedrockruntime"
version = "1.82.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8cb95f77abd4321348dd2f52a25e1de199732f54d2a35860ad20f5df21c66b44"
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",
"once_cell",
"regex-lite",
"tracing",
]
[[package]]
name = "aws-sdk-sqs"
version = "1.64.0"
@@ -918,6 +945,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "69d03c3c05ff80d54ff860fe38c726f6f494c639ae975203a101335f223386db"
dependencies = [
"aws-credential-types",
"aws-smithy-eventstream",
"aws-smithy-http",
"aws-smithy-runtime-api",
"aws-smithy-types",
@@ -945,12 +973,24 @@ dependencies = [
"tokio",
]
[[package]]
name = "aws-smithy-eventstream"
version = "0.60.8"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7c45d3dddac16c5c59d553ece225a88870cf81b7b813c9cc17b78cf4685eac7a"
dependencies = [
"aws-smithy-types",
"bytes",
"crc32fast",
]
[[package]]
name = "aws-smithy-http"
version = "0.62.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c5949124d11e538ca21142d1fba61ab0a2a2c1bc3ed323cdb3e4b878bfb83166"
dependencies = [
"aws-smithy-eventstream",
"aws-smithy-runtime-api",
"aws-smithy-types",
"bytes",
@@ -1257,6 +1297,18 @@ version = "1.7.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "89e25b6adfb930f02d1981565a6e5d9c547ac15a96606256d3b59040e5cd4ca3"
[[package]]
name = "bedrock_test"
version = "1.481.0"
dependencies = [
"aws-config",
"aws-sdk-bedrockruntime",
"serde",
"serde_json",
"tokio",
"tokio-stream",
]
[[package]]
name = "better_scoped_tls"
version = "0.1.2"
@@ -12419,6 +12471,7 @@ dependencies = [
"futures-core",
"pin-project-lite",
"tokio",
"tokio-util",
]
[[package]]
@@ -13730,6 +13783,7 @@ dependencies = [
"async-stream",
"async_zip",
"aws-config",
"aws-sdk-bedrockruntime",
"aws-sdk-sqs",
"axum",
"base32",
@@ -13795,6 +13849,7 @@ dependencies = [
"tokio",
"tokio-native-tls",
"tokio-postgres 0.7.11",
"tokio-stream",
"tokio-tar",
"tokio-tungstenite",
"tokio-util",

View File

@@ -28,7 +28,7 @@ members = [
"./parsers/windmill-parser-py",
"./parsers/windmill-parser-py-imports",
"./parsers/windmill-sql-datatype-parser-wasm",
"./parsers/windmill-parser-yaml", "windmill-macros", "parsers/windmill-parser-nu",
"./parsers/windmill-parser-yaml", "windmill-macros", "parsers/windmill-parser-nu", "bedrock_test",
]
[workspace.package]
@@ -78,6 +78,7 @@ http_trigger = ["windmill-api/http_trigger"]
zip = ["windmill-api/zip"]
static_frontend = ["windmill-api/static_frontend"]
scoped_cache = ["windmill-common/scoped_cache"]
bedrock = ["windmill-api/bedrock"]
# Languages
python = ["windmill-worker/python"]
deno_core = ["windmill-worker/deno_core", "dep:deno_core", "dep:v8"]
@@ -176,6 +177,7 @@ axum = { version = "^0.7", features = ["multipart"] }
headers = "^0"
hyper = { version = "^1", features = ["full"] }
tokio = { version = "^1.42.0", features = ["full", "tracing", "time"] }
tokio-stream = { version = "^0.1", features = ["full"] }
tower = "^0"
tower-http = { version = "^0.6", features = ["trace", "cors"] }
tower-cookies = "^0.10"
@@ -323,9 +325,10 @@ nu-parser = { version = "0.101.0", default-features = false }
datafusion = "39.0.0"
object_store = { version = "0.10.0", features = ["aws", "azure"] }
openidconnect = { version = "4.0.0-rc.1" }
aws-config = "^1"
aws-config = "1.1.6"
aws-sdk-sqs = "1.57.0"
aws-sdk-sts = "^1"
aws-sdk-bedrockruntime = { version = "1.7.0", features = [] }
crc = "^3"
tar = "^0"

View File

@@ -0,0 +1,13 @@
[package]
name = "bedrock_test"
version.workspace = true
authors.workspace = true
edition.workspace = true
[dependencies]
tokio = { version = "1", features = ["full"] }
aws-config = "1.1.6"
aws-sdk-bedrockruntime = "1.82.0"
tokio-stream = "0.1"
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"

View File

@@ -0,0 +1,78 @@
use aws_config::BehaviorVersion;
use aws_sdk_bedrockruntime::{
config::{Credentials, Region},
primitives::Blob,
types::ConverseStreamOutput,
Client as BedrockClient,
};
use serde_json::{json, Value};
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let access_key = "changeme";
let secret_key = "changeme";
let region = "us-east-2";
// ✅ FULL ARN (not short model ID)
let model_id = "arn:aws:bedrock:us-east-2:976079455550:inference-profile/us.anthropic.claude-3-haiku-20240307-v1:0";
let credentials = Credentials::new(access_key, secret_key, None, None, "example");
let config = aws_config::defaults(BehaviorVersion::latest())
.region(Region::new(region.to_string()))
.credentials_provider(credentials)
.load()
.await;
let client = BedrockClient::new(&config);
// ✅ Claude-compatible JSON body
let request_body = json!({
"anthropic_version": "bedrock-2023-05-31",
"messages": [
{
"role": "assistant",
"content": "You are a coding assistant for the Windmill platform. You are provided with a list of `INSTRUCTIONS` and the current contents of a code file under `CODE`. Your task is to respond to the user's request. Assume all user queries are valid and actionable. When the user requests code changes: Always include a single code block with the entire updated file, not just the modified sections. ..."
},
{
"role": "user",
"content": "INSTRUCTIONS:\ntest\n\nWINDMILL LANGUAGE CONTEXT:\nThe user is coding in TypeScript (bun runtime). On Windmill, it is expected that the script exports a single async function called `main`. Do not call the main function. Libraries are installed automatically, do not show how to install them. ..."
}
],
"temperature": 0.0,
"top_p": 0.95,
"max_tokens": 8192,
});
let response = client
.invoke_model_with_response_stream()
.model_id(model_id)
.body(Blob::new(serde_json::to_vec(&request_body)?))
.content_type("application/json")
.send()
.await?;
let mut result = String::new();
let mut stream = response.body;
while let Some(part) = stream.recv().await? {
if let Some(blob) = part.as_chunk().ok().and_then(|c| c.bytes()) {
let bytes = blob.as_ref(); // <- no clone or move
if let Ok(json) = serde_json::from_slice::<Value>(bytes) {
if let Some(text) = json
.get("delta")
.and_then(|d| d.get("text"))
.and_then(|t| t.as_str())
{
print!("{text}");
result.push_str(text);
}
}
}
}
println!("\n\nResponse length: {}", result.len());
Ok(())
}

View File

@@ -32,6 +32,7 @@ postgres_trigger = ["dep:rust-postgres", "dep:pg_escape", "dep:byteorder", "dep:
mqtt_trigger = ["dep:thiserror", "dep:rumqttc"]
sqs_trigger = ["dep:aws-sdk-sqs", "dep:thiserror", "dep:aws-config"]
cloud = ["windmill-common/cloud"]
bedrock = ["dep:aws-config", "dep:aws-sdk-bedrockruntime"]
[dependencies]
windmill-queue.workspace = true
@@ -43,6 +44,7 @@ windmill-parser-py.workspace = true
windmill-git-sync.workspace = true
windmill-indexer = { workspace = true, optional = true }
tokio.workspace = true
tokio-stream.workspace = true
anyhow.workspace = true
argon2.workspace = true
axum.workspace = true
@@ -125,5 +127,6 @@ thiserror = { workspace = true, optional = true }
rust_decimal = { workspace = true, optional = true }
rust-postgres-native-tls = { workspace = true, optional = true}
rumqttc = { workspace = true, optional = true }
aws-sdk-sqs = { workspace = true, optional = true }
aws-config = { workspace = true, optional = true}
aws-sdk-sqs = { workspace = true, optional = true }
aws-config = { workspace = true, optional = true }
aws-sdk-bedrockruntime = { workspace = true, optional = true }

View File

@@ -1997,7 +1997,7 @@ paths:
required:
- premium
- owner
/w/{workspace}/workspaces/threshold_alert:
get:
summary: get threshold alert info
@@ -2368,7 +2368,7 @@ paths:
application/json:
schema:
$ref: "#/components/schemas/AIConfig"
/w/{workspace}/workspaces/edit_error_handler:
post:
@@ -12508,7 +12508,7 @@ components:
AIProvider:
type: string
enum: [openai, azure_openai, anthropic, mistral, deepseek, googleai, groq, openrouter, togetherai, customai]
enum: [openai, azure_openai, anthropic, mistral, deepseek, googleai, groq, openrouter, togetherai, customai, aws_bedrock]
AIProviderModel:
type: object
@@ -13542,7 +13542,7 @@ components:
csharp,
nu,
java
# for related places search: ADD_NEW_LANG
# for related places search: ADD_NEW_LANG
]
Preview:

View File

@@ -10,9 +10,17 @@ use reqwest::{Client, RequestBuilder};
use serde::{Deserialize, Serialize};
use serde_json::value::RawValue;
use std::collections::HashMap;
use tokio_stream::wrappers::ReceiverStream;
use windmill_audit::{audit_ee::audit_log, ActionKind};
use windmill_common::error::{to_anyhow, Error, Result};
#[cfg(feature = "bedrock")]
use aws_sdk_bedrockruntime::{
config::{BehaviorVersion, Credentials, Region},
primitives::Blob,
Client as BedrockClient,
};
lazy_static::lazy_static! {
static ref HTTP_CLIENT: Client = reqwest::ClientBuilder::new()
.timeout(std::time::Duration::from_secs(60 * 5))
@@ -44,6 +52,16 @@ struct AIStandardResource {
organization_id: Option<String>,
}
#[cfg(feature = "bedrock")]
#[derive(Deserialize, Debug)]
struct AIBedrockResource {
region: String,
#[serde(rename = "accessKeyId")]
access_key_id: String,
#[serde(rename = "secretAccessKey")]
secret_access_key: String,
}
#[derive(Deserialize, Debug)]
struct OAuthTokens {
access_token: String,
@@ -52,6 +70,8 @@ struct OAuthTokens {
#[derive(Deserialize, Debug)]
#[serde(untagged)]
enum AIResource {
#[cfg(feature = "bedrock")]
Bedrock(AIBedrockResource),
OAuth(AIOAuthResource),
Standard(AIStandardResource),
}
@@ -63,6 +83,9 @@ struct AIRequestConfig {
pub access_token: Option<String>,
pub organization_id: Option<String>,
pub user: Option<String>,
pub region: Option<String>,
pub access_key_id: Option<String>,
pub secret_access_key: Option<String>,
}
impl AIRequestConfig {
@@ -72,7 +95,28 @@ impl AIRequestConfig {
w_id: &str,
resource: AIResource,
) -> Result<Self> {
let (api_key, access_token, organization_id, base_url, user) = match resource {
tracing::debug!("Creating AI request config for provider: {:?}", provider);
tracing::debug!("Resource: {:?}", resource);
match resource {
#[cfg(feature = "bedrock")]
AIResource::Bedrock(resource) => {
tracing::debug!("Creating Bedrock request config");
let base_url = provider.get_base_url(None, db).await?;
let access_key_id = get_variable_or_self(resource.access_key_id, db, w_id).await?;
let secret_access_key =
get_variable_or_self(resource.secret_access_key, db, w_id).await?;
let region = get_variable_or_self(resource.region, db, w_id).await?;
Ok(Self {
base_url,
access_key_id: Some(access_key_id),
secret_access_key: Some(secret_access_key),
region: Some(region),
organization_id: None,
api_key: None,
access_token: None,
user: None,
})
}
AIResource::Standard(resource) => {
let base_url = provider.get_base_url(resource.base_url, db).await?;
let api_key = if let Some(api_key) = resource.api_key {
@@ -86,7 +130,16 @@ impl AIRequestConfig {
None
};
(api_key, None, organization_id, base_url, None)
Ok(Self {
base_url,
organization_id,
api_key,
access_token: None,
user: None,
region: None,
access_key_id: None,
secret_access_key: None,
})
}
AIResource::OAuth(resource) => {
let user = if let Some(user) = resource.user.clone() {
@@ -97,11 +150,18 @@ impl AIRequestConfig {
let token = Self::get_token_using_oauth(resource, db, w_id).await?;
let base_url = provider.get_base_url(None, db).await?;
(None, Some(token), None, base_url, user)
Ok(Self {
base_url,
organization_id: None,
api_key: None,
access_token: Some(token),
user,
region: None,
access_key_id: None,
secret_access_key: None,
})
}
};
Ok(Self { base_url, organization_id, api_key, access_token, user })
}
}
async fn get_token_using_oauth(
@@ -230,6 +290,9 @@ pub enum AIProvider {
OpenRouter,
TogetherAI,
CustomAI,
#[cfg(feature = "bedrock")]
#[serde(rename = "aws_bedrock")]
AwsBedrock,
}
impl AIProvider {
@@ -265,6 +328,8 @@ 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()),
#[cfg(feature = "bedrock")]
AIProvider::AwsBedrock => Ok("".to_string()), // Bedrock uses AWS SDK directly, not REST API
p @ (AIProvider::CustomAI | AIProvider::AzureOpenAI) => {
if let Some(base_url) = resource_base_url {
Ok(base_url)
@@ -337,6 +402,59 @@ async fn global_proxy(
None => return Err(Error::BadRequest("Provider is required".to_string())),
};
#[cfg(feature = "bedrock")]
if matches!(provider, AIProvider::AwsBedrock) {
let region = headers
.get("X-Bedrock-Region")
.map(|v| v.to_str().unwrap_or("").to_string())
.ok_or_else(|| Error::BadRequest("Bedrock region is required".to_string()))?;
let model_id = headers
.get("X-Bedrock-Model-ID")
.map(|v| v.to_str().unwrap_or("").to_string())
.ok_or_else(|| Error::BadRequest("Bedrock model ID is required".to_string()))?;
// Get AWS credentials from headers if provided
let access_key_id = headers
.get("X-AWS-Access-Key-ID")
.map(|v| v.to_str().unwrap_or("").to_string());
let secret_access_key = headers
.get("X-AWS-Secret-Access-Key")
.map(|v| v.to_str().unwrap_or("").to_string());
tracing::debug!(
"Global Bedrock request - Model: {}, Region: {}, Access Key Provided: {}",
model_id,
region,
access_key_id.is_some()
);
// Use the provided credentials or default AWS credentials from environment
let (status, headers, body) = send_bedrock_request(
body,
Some(&region),
access_key_id.as_deref(),
secret_access_key.as_deref(),
)
.await?;
let mut tx = db.begin().await?;
audit_log(
&mut *tx,
&authed,
"ai.global_request.bedrock",
ActionKind::Execute,
"global",
Some(&authed.email),
None,
)
.await?;
tx.commit().await?;
return Ok((status, headers, body));
}
let Some(api_key) = api_key else {
return Err(Error::BadRequest("API key is required".to_string()));
};
@@ -464,9 +582,28 @@ async fn proxy(
}
};
let request = request_config.prepare_request(&provider, &ai_path, body)?;
let result = if matches!(provider, AIProvider::AwsBedrock) {
send_bedrock_request(
body,
request_config.region.as_deref(),
request_config.access_key_id.as_deref(),
request_config.secret_access_key.as_deref(),
)
.await?
} else {
let request = request_config.prepare_request(&provider, &ai_path, body)?;
let response = request.send().await.map_err(to_anyhow)?;
let response = request.send().await.map_err(to_anyhow)?;
if response.error_for_status_ref().is_err() {
let err_msg = response.text().await.unwrap_or("".to_string());
return Err(Error::AIError(err_msg));
}
let status_code = response.status();
let headers = response.headers().clone();
let stream = response.bytes_stream();
(status_code, headers, axum::body::Body::from_stream(stream))
};
let mut tx = db.begin().await?;
@@ -482,13 +619,209 @@ async fn proxy(
.await?;
tx.commit().await?;
if response.error_for_status_ref().is_err() {
let err_msg = response.text().await.unwrap_or("".to_string());
return Err(Error::AIError(err_msg));
Ok(result)
}
#[cfg(feature = "bedrock")]
async fn create_bedrock_client(
region: Option<&str>,
access_key_id: Option<&str>,
secret_access_key: Option<&str>,
) -> Result<BedrockClient> {
if let (Some(access_key), Some(secret_key), Some(region)) =
(access_key_id, secret_access_key, region)
{
tracing::debug!(
"Creating Bedrock client with provided credentials in region: {}",
region
);
let credentials = Credentials::new(access_key, secret_key, None, None, "DirectTest");
let config = aws_config::defaults(BehaviorVersion::latest())
.region(Region::new(region.to_string()))
.credentials_provider(credentials)
.load()
.await;
Ok(BedrockClient::new(&config))
} else {
tracing::debug!("Creating Bedrock client with default credentials");
let config = aws_config::defaults(BehaviorVersion::latest()).load().await;
Ok(BedrockClient::new(&config))
}
}
#[cfg(feature = "bedrock")]
pub async fn send_bedrock_request(
body: Bytes,
region: Option<&str>,
access_key_id: Option<&str>,
secret_access_key: Option<&str>,
) -> Result<(axum::http::StatusCode, HeaderMap, axum::body::Body)> {
let region = region.ok_or_else(|| Error::BadRequest("Region is required".to_string()))?;
let access_key_id =
access_key_id.ok_or_else(|| Error::BadRequest("Access key ID is required".to_string()))?;
let secret_access_key = secret_access_key
.ok_or_else(|| Error::BadRequest("Secret access key is required".to_string()))?;
tracing::debug!("Sending Bedrock request to region: {}", region);
tracing::debug!("Body: {:#?}", body);
let mut json_body: serde_json::Value = serde_json::from_slice(&body)
.map_err(|e| Error::internal_err(format!("Failed to parse request body: {}", e)))?;
let model_id = json_body["model"]
.as_str()
.ok_or_else(|| Error::internal_err("Model ID not found in request body"))?
.to_string();
// Update anthropic_version default
json_body.as_object_mut().map(|obj| {
obj.entry("anthropic_version")
.or_insert_with(|| serde_json::Value::String("bedrock-2023-05-31".to_string()));
obj.entry("messages")
.or_insert(serde_json::Value::Array(vec![]));
// Add any other required fields with default values
});
// Handle "system" role
if let Some(messages) = json_body.get_mut("messages").and_then(|m| m.as_array_mut()) {
for message in messages {
if let Some(role) = message.get_mut("role") {
// Change "system" role to "user"
if role == "system" {
*role = serde_json::Value::String("user".to_string());
}
}
}
}
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)))
// Remove "tool_choice"
json_body.as_object_mut().map(|obj| {
obj.remove("stream");
obj.remove("model");
obj.remove("tools");
obj.remove("tool_choice");
// Add any other fields you need to remove
});
// Serialize the modified JSON back to bytes
let modified_body = serde_json::to_vec(&json_body)
.map_err(|e| Error::internal_err(format!("Failed to serialize modified body: {}", e)))?;
let client =
create_bedrock_client(Some(region), Some(access_key_id), Some(secret_access_key)).await?;
tracing::debug!("Sending request to Bedrock...");
tracing::debug!("Body: {:#?}", json_body);
let response = client
.invoke_model_with_response_stream()
.model_id(model_id.clone())
.body(Blob::new(modified_body))
.content_type("application/json")
.send()
.await
.map_err(|e| Error::internal_err(format!("Failed to send Bedrock request: {:#?}", e)))?;
let mut stream = response.body;
let (tx, rx): (
tokio::sync::mpsc::Sender<std::result::Result<String, Error>>,
tokio::sync::mpsc::Receiver<std::result::Result<String, Error>>,
) = tokio::sync::mpsc::channel(32);
let tx_clone1 = tx.clone();
let tx_clone2 = tx.clone();
// Spawn a task to process the stream and send chunks
tokio::spawn(async move {
while let Ok(Some(output)) = stream.recv().await {
if let Ok(chunk) = output.as_chunk() {
let chunk_bytes = if let Some(blob) = chunk.bytes() {
blob.as_ref().to_vec()
} else {
tracing::warn!("Received a chunk with no bytes");
continue;
};
// Parse JSON from chunk
if let Ok(json) = serde_json::from_slice::<serde_json::Value>(&chunk_bytes) {
if let Some(text) = json
.get("delta")
.and_then(|d| d.get("text"))
.and_then(|t| t.as_str())
{
let event = serde_json::json!({
"id": "chatcmpl-123",
"object": "chat.completion.chunk",
"created": chrono::Utc::now().timestamp(),
"model": model_id,
"choices": [{
"index": 0,
"delta": {
"content": text
},
"finish_reason": null
}]
});
if let Err(e) = tx_clone1.send(Ok(format!("data: {}\n\n", event.to_string())))
.await
{
tracing::error!("Failed to send chunk: {}", e);
break;
}
}
} else {
tracing::warn!("Non-JSON chunk: {:?}", chunk_bytes);
tracing::debug!("Raw chunk bytes: {:?}", String::from_utf8_lossy(&chunk_bytes));
}
}
}
// Finish event
let final_event = serde_json::json!({
"id": "chatcmpl-123",
"object": "chat.completion.chunk",
"created": chrono::Utc::now().timestamp(),
"model": model_id,
"choices": [{
"index": 0,
"delta": {},
"finish_reason": "stop"
}]
});
let _ = tx_clone1.send(Ok(format!("data: {}\n\n", final_event.to_string())))
.await;
});
// Send heartbeat ping
tokio::spawn(async move {
loop {
tokio::time::sleep(tokio::time::Duration::from_secs(15)).await;
if let Err(e) = tx_clone2.send(Ok(":\n\n".to_string())).await {
tracing::error!("Failed to send heartbeat: {}", e);
break;
}
}
});
// Set up response headers for SSE
let mut response_headers = HeaderMap::new();
response_headers.insert(
http::header::CONTENT_TYPE,
"text/event-stream".parse().unwrap(),
);
response_headers.insert(http::header::CACHE_CONTROL, "no-cache".parse().unwrap());
response_headers.insert(http::header::CONNECTION, "keep-alive".parse().unwrap());
// Create a streaming body from the receiver
let body = axum::body::Body::from_stream(ReceiverStream::new(rx));
Ok((axum::http::StatusCode::OK, response_headers, body))
}

View File

@@ -37,6 +37,7 @@ export const AI_DEFAULT_MODELS: Record<AIProvider, string[]> = {
groq: ['llama-3.3-70b-versatile', 'llama-3.1-8b-instant'],
openrouter: ['meta-llama/llama-3.2-3b-instruct:free'],
togetherai: ['meta-llama/Llama-3.3-70B-Instruct-Turbo'],
aws_bedrock: ['meta.llama3-3-70b-instruct-v1:0','mistral.mistral-7b-instruct-v0:2'],
customai: []
}
@@ -69,6 +70,14 @@ function prepareMessages(aiProvider: AIProvider, messages: ChatCompletionMessage
messages = [...startMessages, ...messages]
}
return messages
case 'aws_bedrock':
// Convert messages to Bedrock format
return messages.map(msg => ({
role: msg.role,
content: typeof msg.content === 'string'
? [{ type: 'text', text: msg.content }]
: msg.content
}))
default:
return messages
}
@@ -96,7 +105,11 @@ 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,
seed: undefined
} as ChatCompletionCreateParams
} as const
class WorkspacedAIClients {

View File

@@ -22,6 +22,7 @@
['groq', 'Groq'],
['openrouter', 'OpenRouter'],
['togetherai', 'Together AI'],
['aws_bedrock', 'Amazon Bedrock'],
['customai', 'Custom AI']
]