use crate::codegen::shared::to_kebab_case;
use crate::codegen::shared::to_pascal_case;
use crate::model::HttpMethod;
use crate::model::ParamSource;
use crate::model::RouterDef;
use crate::model::ToolDef;
use proc_macro2::TokenStream;
use quote::quote;
use syn::spanned::Spanned;
pub fn generate_rest_methods_split(router: &RouterDef) -> (TokenStream, TokenStream) {
let struct_type = &router.struct_type;
let module_name = syn::Ident::new(
&format!(
"__utf_rest_generated_{}",
struct_type
.segments
.last()
.unwrap()
.ident
.to_string()
.to_lowercase()
),
struct_type.span(),
);
let create_router_method = generate_create_rest_router(router, &module_name);
let openapi_method = generate_get_openapi_spec(router);
let param_structs = generate_param_structs(router);
let module_code = quote! {
mod #module_name {
use super::*;
use ::serde::{Serialize, Deserialize};
#param_structs
}
};
let method_code = quote! {
impl #struct_type {
#create_router_method
#openapi_method
}
};
(module_code, method_code)
}
#[allow(dead_code)] pub fn generate_rest_methods(router: &RouterDef) -> TokenStream {
let (module_code, method_code) = generate_rest_methods_split(router);
quote! {
#module_code
#method_code
}
}
fn generate_create_rest_router(router: &RouterDef, module_name: &syn::Ident) -> TokenStream {
let routes = router
.tools
.iter()
.map(|tool| generate_tool_route(tool, module_name));
let prefix: String = router
.metadata
.rest_config
.as_ref()
.and_then(|c| c.prefix.as_ref())
.cloned()
.or_else(|| router.metadata.base_path.clone())
.unwrap_or_else(|| "/api".to_string());
quote! {
pub fn create_rest_router(state: ::std::sync::Arc<Self>) -> ::universal_tool_core::rest::Router {
use ::universal_tool_core::rest::{Router, response::IntoResponse};
let api = Router::new()
#( #routes )*;
Router::new()
.nest(#prefix, api)
.with_state(state)
}
}
}
fn generate_param_structs(router: &RouterDef) -> TokenStream {
router
.tools
.iter()
.map(|tool| {
let struct_name = get_params_struct_name(tool);
let body_params: Vec<_> = tool
.params
.iter()
.filter(|p| matches!(p.source, ParamSource::Body))
.collect();
if body_params.is_empty() {
quote! {}
} else {
let fields = body_params.iter().map(|param| {
let name = ¶m.name;
let ty = ¶m.ty;
let doc = param
.metadata
.description
.as_deref()
.unwrap_or("")
.to_string();
quote! {
#[doc = #doc]
pub #name: #ty
}
});
quote! {
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct #struct_name {
#( #fields ),*
}
}
}
})
.collect()
}
fn generate_tool_route(tool: &ToolDef, module_name: &syn::Ident) -> TokenStream {
let method = determine_http_method(tool);
let path = generate_route_path(tool);
let _handler_name = get_handler_name(tool);
let handler = generate_handler(tool, module_name);
let route_method = match method {
HttpMethod::Get => quote! { ::universal_tool_core::rest::routing::get },
HttpMethod::Post => quote! { ::universal_tool_core::rest::routing::post },
HttpMethod::Put => quote! { ::universal_tool_core::rest::routing::put },
HttpMethod::Delete => quote! { ::universal_tool_core::rest::routing::delete },
HttpMethod::Patch => quote! { ::universal_tool_core::rest::routing::patch },
};
quote! {
.route(#path, #route_method(#handler))
}
}
fn generate_handler(tool: &ToolDef, module_name: &syn::Ident) -> TokenStream {
let has_body_params = tool
.params
.iter()
.any(|p| matches!(p.source, ParamSource::Body));
let param_extractors = generate_param_extractors(tool);
let error_handling = quote! {
let status = match e.code {
::universal_tool_core::error::ErrorCode::BadRequest => ::universal_tool_core::rest::StatusCode::BAD_REQUEST,
::universal_tool_core::error::ErrorCode::InvalidArgument => ::universal_tool_core::rest::StatusCode::BAD_REQUEST,
::universal_tool_core::error::ErrorCode::NotFound => ::universal_tool_core::rest::StatusCode::NOT_FOUND,
::universal_tool_core::error::ErrorCode::PermissionDenied => ::universal_tool_core::rest::StatusCode::FORBIDDEN,
::universal_tool_core::error::ErrorCode::Internal => ::universal_tool_core::rest::StatusCode::INTERNAL_SERVER_ERROR,
::universal_tool_core::error::ErrorCode::Timeout => ::universal_tool_core::rest::StatusCode::REQUEST_TIMEOUT,
::universal_tool_core::error::ErrorCode::Conflict => ::universal_tool_core::rest::StatusCode::CONFLICT,
::universal_tool_core::error::ErrorCode::NetworkError => ::universal_tool_core::rest::StatusCode::SERVICE_UNAVAILABLE,
::universal_tool_core::error::ErrorCode::ExternalServiceError => ::universal_tool_core::rest::StatusCode::BAD_GATEWAY,
::universal_tool_core::error::ErrorCode::ExecutionFailed => ::universal_tool_core::rest::StatusCode::INTERNAL_SERVER_ERROR,
::universal_tool_core::error::ErrorCode::SerializationError => ::universal_tool_core::rest::StatusCode::BAD_REQUEST,
::universal_tool_core::error::ErrorCode::IoError => ::universal_tool_core::rest::StatusCode::INTERNAL_SERVER_ERROR,
};
(status, ::universal_tool_core::rest::Json(::serde_json::json!({
"error": e.to_string(),
"code": format!("{:?}", e.code),
}))).into_response()
};
let method_args_vec: Vec<TokenStream> = tool
.params
.iter()
.map(|param| {
let name = ¶m.name;
match param.source {
ParamSource::Body => quote! { params.#name },
_ => quote! { #name },
}
})
.collect();
let method_call = crate::codegen::shared::generate_normalized_method_call(
tool,
quote! { state },
method_args_vec,
);
if has_body_params {
let params_struct = get_params_struct_name(tool);
quote! {
|::universal_tool_core::rest::State(state): ::universal_tool_core::rest::State<::std::sync::Arc<Self>>#param_extractors,
::universal_tool_core::rest::Json(params): ::universal_tool_core::rest::Json<#module_name::#params_struct>| async move {
match #method_call {
Ok(result) => (::universal_tool_core::rest::StatusCode::OK, ::universal_tool_core::rest::Json(result)).into_response(),
Err(e) => { #error_handling },
}
}
}
} else {
quote! {
|::universal_tool_core::rest::State(state): ::universal_tool_core::rest::State<::std::sync::Arc<Self>>#param_extractors| async move {
match #method_call {
Ok(result) => (::universal_tool_core::rest::StatusCode::OK, ::universal_tool_core::rest::Json(result)).into_response(),
Err(e) => { #error_handling },
}
}
}
}
}
fn generate_param_extractors(tool: &ToolDef) -> TokenStream {
let extractors = tool.params.iter()
.filter(|p| !matches!(p.source, ParamSource::Body))
.map(|param| {
let name = ¶m.name;
let ty = ¶m.ty;
match param.source {
ParamSource::Path => quote! { , ::universal_tool_core::rest::Path(#name): ::universal_tool_core::rest::Path<#ty> },
ParamSource::Query => quote! { , ::universal_tool_core::rest::Query(#name): ::universal_tool_core::rest::Query<#ty> },
_ => quote! {},
}
});
quote! { #( #extractors )* }
}
fn generate_get_openapi_spec(_router: &RouterDef) -> TokenStream {
quote! {
pub fn get_openapi_spec(&self) -> String {
"OpenAPI generation requires the 'openapi' feature to be enabled".to_string()
}
}
}
fn determine_http_method(tool: &ToolDef) -> HttpMethod {
if let Some(rest_config) = &tool.metadata.rest_config {
return rest_config.method;
}
let name = tool.method_name.to_string();
if name.starts_with("get") || name.starts_with("list") || name.starts_with("find") {
HttpMethod::Get
} else if name.starts_with("create") || name.starts_with("add") || name.starts_with("new") {
HttpMethod::Post
} else if name.starts_with("update") || name.starts_with("modify") || name.starts_with("set") {
HttpMethod::Put
} else if name.starts_with("delete") || name.starts_with("remove") {
HttpMethod::Delete
} else if name.starts_with("patch") {
HttpMethod::Patch
} else {
HttpMethod::Post
}
}
fn generate_route_path(tool: &ToolDef) -> String {
if let Some(rest_config) = &tool.metadata.rest_config
&& let Some(path) = &rest_config.path
{
return path.clone();
}
let base_path = format!("/{}", to_kebab_case(&tool.tool_name));
let mut path = base_path;
for param in &tool.params {
if matches!(param.source, ParamSource::Path) {
path.push_str(&format!("/:{}", param.name));
}
}
path
}
fn get_params_struct_name(tool: &ToolDef) -> syn::Ident {
let name = format!("{}Params", to_pascal_case(&tool.tool_name));
syn::Ident::new(&name, tool.method_name.span())
}
fn get_handler_name(tool: &ToolDef) -> syn::Ident {
let name = format!("handle_rest_{}", tool.method_name);
syn::Ident::new(&name, tool.method_name.span())
}
#[allow(dead_code)] fn generate_openapi_schemas(router: &RouterDef) -> TokenStream {
let schemas = router
.tools
.iter()
.filter(|tool| {
tool.params
.iter()
.any(|p| matches!(p.source, ParamSource::Body))
})
.map(|_tool| {
quote! {
}
});
quote! { #( #schemas )* }
}
#[allow(dead_code)] fn generate_openapi_paths(router: &RouterDef) -> TokenStream {
let paths = router.tools.iter().map(|tool| {
let fn_name = syn::Ident::new(
&format!("__openapi_path_{}", tool.method_name),
tool.method_name.span()
);
let path = generate_route_path(tool);
let method = determine_http_method(tool);
let method_str = format!("{method:?}").to_lowercase();
let description = &tool.metadata.description;
let has_body = tool.params.iter().any(|p| matches!(p.source, ParamSource::Body));
let request_body = if has_body {
let struct_name = get_params_struct_name(tool);
quote! { request_body = #struct_name, }
} else {
quote! {}
};
let return_type = &tool.return_type;
quote! {
#[::utoipa::path(
#method_str,
path = #path,
#request_body
responses(
(status = 200, description = #description, body = #return_type),
(status = 400, description = "Bad request", body = ::universal_tool_core::error::ToolError),
(status = 500, description = "Internal server error", body = ::universal_tool_core::error::ToolError)
)
)]
async fn #fn_name() {}
}
});
quote! { #( #paths )* }
}