use crate::{
Api, HttpMethod, Operation, OperationMediaType, SecuritySchemeCatalog, SecuritySchemeKind,
};
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct SdkSemantics {
pub operations: Vec<OperationSemantics>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct OperationSemantics {
pub operation_id: String,
pub auth: Vec<AuthAlternative>,
pub errors: Vec<DeclaredError>,
pub retry: RetryClass,
pub streaming: Option<StreamingKind>,
pub request_body: Option<RequestBodyKind>,
pub pagination: Option<PaginationHint>,
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct AuthAlternative {
pub schemes: Vec<AuthScheme>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum AuthScheme {
ApiKey { name: String, location: String },
Http { scheme: String },
OAuth2 { scopes: Vec<String> },
OpenIdConnect { discovery_url: Option<String> },
Other { name: String },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct DeclaredError {
pub status: String,
pub body_type: Option<String>,
pub content_type: Option<String>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RetryClass {
Idempotent,
IdempotencyKey,
Unsafe,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum StreamingKind {
ServerSentEvents,
Binary,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RequestBodyKind {
Json,
Multipart,
FormUrlEncoded,
Binary,
Other,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PaginationHint {
pub source: PaginationSource,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum PaginationSource {
Poolster,
Speakeasy,
}
pub fn analyze_sdk_semantics(
api: &Api,
security_schemes: Option<&SecuritySchemeCatalog>,
) -> SdkSemantics {
SdkSemantics {
operations: api
.operations
.iter()
.map(|operation| analyze_operation(operation, security_schemes))
.collect(),
}
}
pub fn analyze_operation(
operation: &Operation,
security_schemes: Option<&SecuritySchemeCatalog>,
) -> OperationSemantics {
OperationSemantics {
operation_id: operation.id.clone(),
auth: operation
.security
.iter()
.map(|requirement| AuthAlternative {
schemes: requirement
.schemes
.iter()
.map(|(name, scopes)| resolve_auth(name, scopes, security_schemes))
.collect(),
})
.collect(),
errors: operation
.responses
.iter()
.filter(|response| is_error_status(&response.status))
.map(|response| DeclaredError {
status: response.status.clone(),
body_type: response
.media_types
.first()
.and_then(|media| media.schema.as_ref())
.and_then(|schema| schema.kind.reference_name())
.map(str::to_owned),
content_type: response
.media_types
.first()
.map(|media| media.content_type.clone()),
})
.collect(),
retry: retry_class(operation),
streaming: streaming_kind(operation),
request_body: operation
.request_body
.as_ref()
.and_then(|body| body.media_types.first())
.map(request_body_kind),
pagination: pagination_hint(operation),
}
}
fn resolve_auth(
name: &str,
scopes: &[String],
security_schemes: Option<&SecuritySchemeCatalog>,
) -> AuthScheme {
let scheme = security_schemes
.and_then(|catalog| catalog.schemes.iter().find(|scheme| scheme.name == name));
match scheme.map(|scheme| &scheme.kind) {
Some(SecuritySchemeKind::ApiKey { name, location }) => AuthScheme::ApiKey {
name: name.clone().unwrap_or_else(|| "Authorization".into()),
location: location.clone().unwrap_or_else(|| "header".into()),
},
Some(SecuritySchemeKind::Http { scheme, .. }) => AuthScheme::Http {
scheme: scheme.clone().unwrap_or_else(|| "bearer".into()),
},
Some(SecuritySchemeKind::OAuth2 { .. }) => AuthScheme::OAuth2 {
scopes: scopes.to_vec(),
},
Some(SecuritySchemeKind::OpenIdConnect { discovery_url }) => AuthScheme::OpenIdConnect {
discovery_url: discovery_url.clone(),
},
Some(SecuritySchemeKind::Other { .. }) | None => AuthScheme::Other { name: name.into() },
}
}
fn is_error_status(status: &str) -> bool {
status == "default"
|| status
.parse::<u16>()
.is_ok_and(|status| (400..600).contains(&status))
}
fn retry_class(operation: &Operation) -> RetryClass {
let resolved_header = operation
.annotations
.get(crate::idempotency::RESOLVED_ANNOTATION)
.and_then(|value| value.get("header"))
.and_then(serde_json::Value::as_str);
if operation.parameters.iter().any(|parameter| {
parameter.location == "header"
&& (parameter.name.eq_ignore_ascii_case("idempotency-key")
|| resolved_header
.is_some_and(|header| parameter.name.eq_ignore_ascii_case(header)))
}) {
return RetryClass::IdempotencyKey;
}
match operation.method {
HttpMethod::Get
| HttpMethod::Put
| HttpMethod::Delete
| HttpMethod::Head
| HttpMethod::Options
| HttpMethod::Trace
| HttpMethod::Query => RetryClass::Idempotent,
HttpMethod::Post | HttpMethod::Patch | HttpMethod::Custom(_) => RetryClass::Unsafe,
}
}
fn streaming_kind(operation: &Operation) -> Option<StreamingKind> {
operation
.responses
.iter()
.flat_map(|response| &response.media_types)
.find_map(|media| {
let content_type = media
.content_type
.split(';')
.next()
.unwrap_or_default()
.trim()
.to_ascii_lowercase();
match content_type.as_str() {
"text/event-stream" => Some(StreamingKind::ServerSentEvents),
"application/octet-stream" | "application/pdf" | "image/png" | "image/jpeg" => {
Some(StreamingKind::Binary)
}
_ if content_type.starts_with("multipart/") => Some(StreamingKind::Binary),
_ => None,
}
})
}
fn request_body_kind(media: &OperationMediaType) -> RequestBodyKind {
let content_type = media
.content_type
.split(';')
.next()
.unwrap_or_default()
.trim()
.to_ascii_lowercase();
match content_type.as_str() {
"application/json" | "application/problem+json" => RequestBodyKind::Json,
_ if content_type.starts_with("multipart/") => RequestBodyKind::Multipart,
"application/x-www-form-urlencoded" => RequestBodyKind::FormUrlEncoded,
"application/octet-stream" => RequestBodyKind::Binary,
_ => RequestBodyKind::Other,
}
}
fn pagination_hint(operation: &Operation) -> Option<PaginationHint> {
if operation.annotations.contains_key("x-poolster-pagination") {
Some(PaginationHint {
source: PaginationSource::Poolster,
})
} else if operation.annotations.contains_key("x-speakeasy-pagination") {
Some(PaginationHint {
source: PaginationSource::Speakeasy,
})
} else {
None
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use serde_json::json;
use crate::{
OperationMediaType, OperationParameter, OperationResponse, SchemaKind, SchemaValue,
SecurityRequirement, SecurityScheme,
};
use super::*;
#[test]
fn query_and_standard_safe_methods_are_replayable() {
for method in [
HttpMethod::Query,
HttpMethod::Head,
HttpMethod::Options,
HttpMethod::Trace,
] {
assert_eq!(
retry_class(&Operation {
method,
..Operation::default()
}),
RetryClass::Idempotent
);
}
}
#[test]
fn patch_requires_explicit_key_and_custom_resolved_header_is_recognized() {
let mut operation = Operation {
id: "patchItem".into(),
method: HttpMethod::Patch,
..Default::default()
};
assert_eq!(retry_class(&operation), RetryClass::Unsafe);
let mut api = Api::default();
operation.annotations.insert(
"x-poolster-idempotency".into(),
serde_json::json!({"header":"X-Request-Key"}),
);
api.operations.push(operation);
let prepared = crate::idempotency::prepare_api(&api, &Default::default()).unwrap();
assert_eq!(
retry_class(&prepared.operations[0]),
RetryClass::IdempotencyKey
);
}
#[test]
fn resolves_auth_errors_retry_and_media_from_one_operation() {
let operation = Operation {
id: "createMessage".into(),
method: HttpMethod::Post,
path: "/messages".into(),
parameters: vec![OperationParameter {
name: "Idempotency-Key".into(),
location: "header".into(),
required: false,
schema: None,
description: None,
annotations: BTreeMap::new(),
}],
request_body: Some(crate::OperationRequestBody {
required: true,
description: None,
media_types: vec![OperationMediaType {
content_type: "multipart/form-data".into(),
schema: None,
}],
}),
responses: vec![
OperationResponse {
status: "201".into(),
description: None,
media_types: vec![],
},
OperationResponse {
status: "429".into(),
description: None,
media_types: vec![OperationMediaType {
content_type: "application/json".into(),
schema: Some(SchemaValue::reference(
"#/components/schemas/RateLimitError",
)),
}],
},
OperationResponse {
status: "default".into(),
description: None,
media_types: vec![],
},
],
security: vec![SecurityRequirement {
schemes: BTreeMap::from([("apiKey".into(), vec![])]),
}],
annotations: BTreeMap::from([("x-poolster-pagination".into(), json!({}))]),
};
let catalog = SecuritySchemeCatalog {
schemes: vec![SecurityScheme {
name: "apiKey".into(),
description: None,
kind: SecuritySchemeKind::ApiKey {
name: Some("X-API-Key".into()),
location: Some("header".into()),
},
}],
};
let semantics = analyze_operation(&operation, Some(&catalog));
assert_eq!(semantics.retry, RetryClass::IdempotencyKey);
assert_eq!(semantics.request_body, Some(RequestBodyKind::Multipart));
assert_eq!(semantics.errors.len(), 2);
assert_eq!(
semantics.errors[0].body_type.as_deref(),
Some("RateLimitError")
);
assert_eq!(
semantics.auth,
vec![AuthAlternative {
schemes: vec![AuthScheme::ApiKey {
name: "X-API-Key".into(),
location: "header".into(),
}],
}]
);
assert_eq!(
semantics.pagination,
Some(PaginationHint {
source: PaginationSource::Poolster,
})
);
}
#[test]
fn detects_server_sent_events_and_does_not_guess_pagination() {
let operation = Operation {
id: "watchEvents".into(),
responses: vec![OperationResponse {
status: "200".into(),
description: None,
media_types: vec![OperationMediaType {
content_type: "text/event-stream".into(),
schema: Some(SchemaValue::new(SchemaKind::String)),
}],
}],
..Operation::default()
};
let semantics = analyze_operation(&operation, None);
assert_eq!(semantics.streaming, Some(StreamingKind::ServerSentEvents));
assert_eq!(semantics.pagination, None);
assert_eq!(semantics.retry, RetryClass::Idempotent);
}
}
#[cfg(test)]
mod custom_method_tests {
use super::*;
#[test]
fn custom_method_is_unsafe_without_explicit_idempotency() {
let operation = Operation {
method: HttpMethod::Custom("COPY".into()),
..Default::default()
};
assert_eq!(retry_class(&operation), RetryClass::Unsafe);
}
}