refactor: extract provider pool abstractions

This commit is contained in:
fawney19
2026-05-13 18:19:15 +08:00
parent 3c2497f019
commit 5d1460e051
55 changed files with 3469 additions and 2184 deletions

View File

@@ -11,6 +11,10 @@ use sha2::{Digest, Sha256};
use std::collections::BTreeMap;
pub const KIRO_PROVIDER_TYPE: &str = "kiro";
pub const DEFAULT_REGION: &str = "us-east-1";
pub const DEFAULT_KIRO_VERSION: &str = "0.3.210";
pub const DEFAULT_NODE_VERSION: &str = "22.21.1";
pub const DEFAULT_SYSTEM_VERSION: &str = "other#unknown";
const IDC_AMZ_USER_AGENT: &str =
"aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE";
@@ -39,7 +43,8 @@ impl KiroAuthConfig {
auth_method: string_field(
object,
&["auth_method", "authMethod", "auth_type", "authType"],
),
)
.map(|value| normalize_kiro_auth_method(&value)),
refresh_token: string_field(object, &["refresh_token", "refreshToken"]),
expires_at: u64_field(object.get("expires_at"))
.or_else(|| u64_field(object.get("expiresAt"))),
@@ -93,7 +98,7 @@ impl KiroAuthConfig {
.or(self.region.as_deref())
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or("us-east-1")
.unwrap_or(DEFAULT_REGION)
}
pub fn effective_api_region(&self) -> &str {
@@ -101,7 +106,7 @@ impl KiroAuthConfig {
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or("us-east-1")
.unwrap_or(DEFAULT_REGION)
}
pub fn effective_kiro_version(&self) -> &str {
@@ -109,40 +114,111 @@ impl KiroAuthConfig {
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or("0.3.210")
.unwrap_or(DEFAULT_KIRO_VERSION)
}
pub fn effective_system_version(&self) -> &str {
self.system_version
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(DEFAULT_SYSTEM_VERSION)
}
pub fn effective_node_version(&self) -> &str {
self.node_version
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(DEFAULT_NODE_VERSION)
}
pub fn cached_access_token(&self) -> Option<&str> {
self.access_token
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
}
pub fn cached_access_token_requires_refresh(&self, skew_seconds: u64) -> bool {
let Some(expires_at) = self.expires_at else {
return self.can_refresh_access_token();
};
let now = current_unix_secs();
now >= expires_at.saturating_sub(skew_seconds)
}
pub fn is_idc_auth(&self) -> bool {
self.auth_method
let explicit_method = self
.auth_method
.as_deref()
.map(normalize_kiro_auth_method)
.unwrap_or_else(|| "social".to_string());
if explicit_method != "social" {
return matches!(explicit_method.as_str(), "idc" | "external_idp");
}
self.client_id
.as_deref()
.map(str::trim)
.map(str::to_ascii_lowercase)
.is_some_and(|value| matches!(value.as_str(), "idc" | "external_idp"))
|| (self
.client_id
.filter(|value| !value.is_empty())
.is_some()
&& self
.client_secret
.as_deref()
.is_some_and(|value| !value.trim().is_empty())
&& self
.client_secret
.as_deref()
.is_some_and(|value| !value.trim().is_empty()))
.map(str::trim)
.filter(|value| !value.is_empty())
.is_some()
}
pub fn uses_external_idp_token_type(&self) -> bool {
self.auth_method
.as_deref()
.map(normalize_kiro_auth_method)
.as_deref()
== Some("external_idp")
}
pub fn profile_arn_for_payload(&self) -> Option<&str> {
if self.is_idc_auth() {
return None;
}
self.profile_arn
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
}
pub fn profile_arn_for_mcp(&self) -> Option<&str> {
self.profile_arn
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
}
pub fn can_refresh_access_token(&self) -> bool {
self.refresh_token
let refresh_token = self
.refresh_token
.as_deref()
.map(str::trim)
.filter(|value| value.len() >= 100 && !value.contains("..."))
.filter(|value| !value.is_empty())
.filter(|value| value.len() >= 100 && !value.contains("..."));
if refresh_token.is_none() {
return false;
}
if !self.is_idc_auth() {
return true;
}
self.client_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.is_some()
&& (!self.is_idc_auth()
|| (self
.client_id
.as_deref()
.is_some_and(|value| !value.trim().is_empty())
&& self
.client_secret
.as_deref()
.is_some_and(|value| !value.trim().is_empty())))
&& self
.client_secret
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.is_some()
}
}
@@ -163,7 +239,7 @@ impl KiroProviderOAuthAdapter {
self
}
async fn refresh_auth_config(
pub async fn refresh_auth_config(
&self,
executor: &dyn OAuthHttpExecutor,
ctx: &ProviderOAuthTransportContext,
@@ -436,7 +512,7 @@ pub fn generate_kiro_machine_id(
if let Some(machine_id) = auth_config
.machine_id
.as_deref()
.and_then(normalize_machine_id)
.and_then(normalize_kiro_machine_id)
{
return Some(machine_id);
}
@@ -497,7 +573,7 @@ fn resolve_expires_at(payload: &Value) -> u64 {
current_unix_secs().saturating_add(expires_in)
}
fn normalize_machine_id(raw: &str) -> Option<String> {
pub fn normalize_kiro_machine_id(raw: &str) -> Option<String> {
let raw = raw.trim();
if raw.len() == 64 && raw.bytes().all(|byte| byte.is_ascii_hexdigit()) {
return Some(raw.to_ascii_lowercase());
@@ -514,6 +590,26 @@ fn normalize_machine_id(raw: &str) -> Option<String> {
None
}
fn normalize_kiro_auth_method(raw: &str) -> String {
let value = raw.trim().to_ascii_lowercase();
match value.as_str() {
"" => "social".to_string(),
"builder-id"
| "builder_id"
| "builderid"
| "device"
| "device-auth"
| "device_authorization"
| "iam"
| "identity-center"
| "identity_center"
| "identitycenter"
| "idc" => "idc".to_string(),
"external-idp" | "external_idp" | "externalidp" => "external_idp".to_string(),
_ => value,
}
}
fn string_field(object: &serde_json::Map<String, Value>, keys: &[&str]) -> Option<String> {
keys.iter()
.find_map(|key| object.get(*key))
@@ -548,7 +644,51 @@ fn secret_fingerprint(value: &str) -> String {
#[cfg(test)]
mod tests {
use super::{generate_kiro_machine_id, KiroAuthConfig};
use super::{
generate_kiro_machine_id, KiroAuthConfig, KiroProviderOAuthAdapter, IDC_AMZ_USER_AGENT,
};
use crate::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse};
use crate::provider::ProviderOAuthTransportContext;
use async_trait::async_trait;
use serde_json::json;
use std::sync::{Arc, Mutex};
#[derive(Debug, Clone)]
struct StaticExecutor {
seen_request: Arc<Mutex<Option<OAuthHttpRequest>>>,
response: serde_json::Value,
}
#[async_trait]
impl OAuthHttpExecutor for StaticExecutor {
async fn execute(
&self,
request: OAuthHttpRequest,
) -> Result<OAuthHttpResponse, crate::core::OAuthError> {
*self.seen_request.lock().expect("mutex should lock") = Some(request);
Ok(OAuthHttpResponse {
status_code: 200,
body_text: self.response.to_string(),
json_body: Some(self.response.clone()),
})
}
}
fn test_ctx() -> ProviderOAuthTransportContext {
ProviderOAuthTransportContext {
provider_id: "provider-1".to_string(),
provider_type: "kiro".to_string(),
endpoint_id: None,
key_id: Some("key-1".to_string()),
auth_type: Some("oauth".to_string()),
decrypted_api_key: None,
decrypted_auth_config: None,
provider_config: None,
endpoint_config: None,
key_config: None,
network: crate::network::OAuthNetworkContext::provider_operation(None),
}
}
#[test]
fn normalizes_kiro_uuid_machine_id() {
@@ -573,4 +713,151 @@ mod tests {
Some("123e4567e89b12d3a456426614174000123e4567e89b12d3a456426614174000")
);
}
#[tokio::test]
async fn refreshes_social_auth_config_with_provider_adapter() {
let seen_request = Arc::new(Mutex::new(None));
let executor = StaticExecutor {
seen_request: Arc::clone(&seen_request),
response: json!({
"accessToken": "new-social-access-token",
"refreshToken": "s".repeat(120),
"expiresIn": 3600,
"profileArn": "arn:aws:bedrock:demo"
}),
};
let adapter = KiroProviderOAuthAdapter::default()
.with_refresh_base_urls(Some("https://auth.example.test".to_string()), None);
let auth_config = KiroAuthConfig {
auth_method: Some("social".to_string()),
refresh_token: Some("r".repeat(120)),
expires_at: Some(1),
profile_arn: None,
region: None,
auth_region: None,
api_region: None,
client_id: None,
client_secret: None,
machine_id: Some("123e4567-e89b-12d3-a456-426614174000".to_string()),
kiro_version: Some("1.2.3".to_string()),
system_version: None,
node_version: None,
access_token: None,
};
let refreshed = adapter
.refresh_auth_config(&executor, &test_ctx(), &auth_config)
.await
.expect("social refresh should succeed");
assert_eq!(
refreshed.access_token.as_deref(),
Some("new-social-access-token")
);
let expected_rotated_refresh_token = "s".repeat(120);
assert_eq!(
refreshed.refresh_token.as_deref(),
Some(expected_rotated_refresh_token.as_str())
);
assert_eq!(
refreshed.profile_arn.as_deref(),
Some("arn:aws:bedrock:demo")
);
assert!(refreshed.expires_at.is_some());
let seen = seen_request
.lock()
.expect("mutex should lock")
.clone()
.expect("request should be captured");
assert_eq!(seen.request_id, "provider-oauth:kiro-social-refresh");
assert_eq!(seen.method, reqwest::Method::POST);
assert_eq!(seen.url, "https://auth.example.test/refreshToken");
assert_eq!(
seen.headers.get("user-agent").map(String::as_str),
Some("KiroIDE-1.2.3-123e4567e89b12d3a456426614174000123e4567e89b12d3a456426614174000")
);
let expected_refresh_token = "r".repeat(120);
assert_eq!(
seen.json_body
.as_ref()
.and_then(|body| body.get("refreshToken"))
.and_then(serde_json::Value::as_str),
Some(expected_refresh_token.as_str())
);
}
#[tokio::test]
async fn refreshes_idc_auth_config_with_provider_adapter() {
let seen_request = Arc::new(Mutex::new(None));
let executor = StaticExecutor {
seen_request: Arc::clone(&seen_request),
response: json!({
"accessToken": "new-idc-access-token",
"refreshToken": "i".repeat(120),
"expiresIn": 1800
}),
};
let adapter = KiroProviderOAuthAdapter::default()
.with_refresh_base_urls(None, Some("https://idc.example.test".to_string()));
let auth_config = KiroAuthConfig {
auth_method: Some("identity_center".to_string()),
refresh_token: Some("r".repeat(120)),
expires_at: Some(1),
profile_arn: Some("arn:aws:bedrock:demo".to_string()),
region: None,
auth_region: None,
api_region: None,
client_id: Some("client-id".to_string()),
client_secret: Some("client-secret".to_string()),
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: None,
};
let refreshed = adapter
.refresh_auth_config(&executor, &test_ctx(), &auth_config)
.await
.expect("idc refresh should succeed");
assert_eq!(
refreshed.access_token.as_deref(),
Some("new-idc-access-token")
);
let expected_rotated_refresh_token = "i".repeat(120);
assert_eq!(
refreshed.refresh_token.as_deref(),
Some(expected_rotated_refresh_token.as_str())
);
assert!(refreshed.machine_id.is_some());
assert!(refreshed.expires_at.is_some());
let seen = seen_request
.lock()
.expect("mutex should lock")
.clone()
.expect("request should be captured");
assert_eq!(seen.request_id, "provider-oauth:kiro-idc-refresh");
assert_eq!(seen.method, reqwest::Method::POST);
assert_eq!(seen.url, "https://idc.example.test/token");
assert_eq!(
seen.headers.get("user-agent").map(String::as_str),
Some("node")
);
assert_eq!(
seen.headers.get("x-amz-user-agent").map(String::as_str),
Some(IDC_AMZ_USER_AGENT)
);
let body = seen.json_body.expect("json body should exist");
assert_eq!(body.get("grantType"), Some(&json!("refresh_token")));
assert_eq!(body.get("clientId"), Some(&json!("client-id")));
assert_eq!(body.get("clientSecret"), Some(&json!("client-secret")));
let expected_refresh_token = "r".repeat(120);
assert_eq!(
body.get("refreshToken").and_then(serde_json::Value::as_str),
Some(expected_refresh_token.as_str())
);
}
}

View File

@@ -9,5 +9,7 @@ pub use generic::{
GenericProviderOAuthAdapter, GenericProviderOAuthTemplate, GENERIC_PROVIDER_OAUTH_TEMPLATES,
};
pub use kiro::{
generate_kiro_machine_id, KiroAuthConfig, KiroProviderOAuthAdapter, KIRO_PROVIDER_TYPE,
generate_kiro_machine_id, normalize_kiro_machine_id, KiroAuthConfig, KiroProviderOAuthAdapter,
DEFAULT_KIRO_VERSION, DEFAULT_NODE_VERSION, DEFAULT_REGION, DEFAULT_SYSTEM_VERSION,
KIRO_PROVIDER_TYPE,
};