use crate::ir::{IrResponse, IrReturnType, IrSseReturn, IrType};
use crate::parse::media_type::MediaType;
use crate::parse::response::ResponseOrRef;
use crate::parse::schema::SchemaOrRef;
use super::name_normalizer::normalize_name;
use super::schema_resolver::schema_or_ref_to_ir_type;
use indexmap::IndexMap;
pub fn detect_return_type(
operation_id: &str,
responses: &IndexMap<String, ResponseOrRef>,
) -> IrReturnType {
let success_response = find_success_response(responses);
let Some((status_code, response)) = success_response else {
return IrReturnType::Void;
};
let content = match response {
ResponseOrRef::Response(r) => &r.content,
ResponseOrRef::Ref { .. } => return IrReturnType::Void,
};
if content.is_empty() {
return IrReturnType::Void;
}
let sse = content.get("text/event-stream");
let json = content.get("application/json");
match (sse, json) {
(Some(sse_mt), json_mt) => {
let sse_return = build_sse_return(operation_id, sse_mt, json_mt, status_code);
IrReturnType::Sse(sse_return)
}
(None, Some(json_mt)) => {
let response_type = match &json_mt.schema {
Some(s) => schema_or_ref_to_ir_type(s),
None => IrType::Any,
};
let description = match response {
ResponseOrRef::Response(r) => Some(r.description.clone()),
_ => None,
};
IrReturnType::Standard(IrResponse {
response_type,
description,
status_code,
})
}
(None, None) => {
if let Some((_ct, mt)) = content.first() {
let response_type = match &mt.schema {
Some(s) => schema_or_ref_to_ir_type(s),
None => IrType::Any,
};
IrReturnType::Standard(IrResponse {
response_type,
description: None,
status_code,
})
} else {
IrReturnType::Void
}
}
}
}
fn build_sse_return(
operation_id: &str,
sse_mt: &MediaType,
json_mt: Option<&MediaType>,
status_code: u16,
) -> IrSseReturn {
let (event_type, variants, event_type_name) = match &sse_mt.item_schema {
Some(item_schema) => extract_event_info(operation_id, item_schema),
None => {
match &sse_mt.schema {
Some(s) => (schema_or_ref_to_ir_type(s), vec![], None),
None => (IrType::Any, vec![], None),
}
}
};
let json_response = json_mt.map(|mt| {
let response_type = match &mt.schema {
Some(s) => schema_or_ref_to_ir_type(s),
None => IrType::Any,
};
IrResponse {
response_type,
description: None,
status_code,
}
});
IrSseReturn {
event_type,
variants,
event_type_name,
also_has_json: json_response.is_some(),
json_response,
}
}
fn extract_event_info(
operation_id: &str,
item_schema: &SchemaOrRef,
) -> (IrType, Vec<IrType>, Option<String>) {
match item_schema {
SchemaOrRef::Ref { .. } => {
let ir_type = schema_or_ref_to_ir_type(item_schema);
(ir_type, vec![], None)
}
SchemaOrRef::Schema(schema) => {
if !schema.one_of.is_empty() {
let variants: Vec<IrType> =
schema.one_of.iter().map(schema_or_ref_to_ir_type).collect();
let event_name = format!("{}StreamEvent", normalize_name(operation_id).pascal_case);
let event_type = IrType::Union(variants.clone());
(event_type, variants, Some(event_name))
} else {
let ir_type = schema_or_ref_to_ir_type(item_schema);
(ir_type, vec![], None)
}
}
}
}
fn find_success_response(
responses: &IndexMap<String, ResponseOrRef>,
) -> Option<(u16, &ResponseOrRef)> {
const PREFERRED: [&str; 4] = ["200", "201", "202", "203"];
let ordered = PREFERRED
.iter()
.filter_map(|code| responses.get_key_value(*code))
.chain(
responses
.iter()
.filter(|(code, _)| is_2xx(code) && !PREFERRED.contains(&code.as_str())),
)
.chain(responses.get_key_value("2XX"))
.chain(responses.get_key_value("2xx"))
.chain(responses.get_key_value("default"));
let mut first: Option<(&String, &ResponseOrRef)> = None;
for (code, response) in ordered {
if first.is_none() {
first = Some((code, response));
}
if has_body(response) {
return Some((status_from_code(code), response));
}
}
first.map(|(code, response)| (status_from_code(code), response))
}
fn status_from_code(code: &str) -> u16 {
code.parse().unwrap_or(200)
}
fn has_body(response: &ResponseOrRef) -> bool {
match response {
ResponseOrRef::Response(r) => !r.content.is_empty(),
ResponseOrRef::Ref { .. } => false,
}
}
fn is_2xx(code: &str) -> bool {
code.len() == 3 && code.starts_with('2') && code.as_bytes()[1..].iter().all(u8::is_ascii_digit)
}