Skip to main content

salvo_oapi/
naming.rs

1use std::any::TypeId;
2use std::collections::BTreeMap;
3use std::sync::LazyLock;
4
5use parking_lot::{RwLock, RwLockReadGuard};
6use regex::Regex;
7
8/// NameRule is used to specify the rule of naming.
9#[derive(Default, Debug, Clone, Copy)]
10pub enum NameRule {
11    /// Auto generate name by namer.
12    #[default]
13    Auto,
14    /// Force to use the given name.
15    Force(&'static str),
16}
17
18static GLOBAL_NAMER: LazyLock<RwLock<Box<dyn Namer>>> =
19    LazyLock::new(|| RwLock::new(Box::new(FlexNamer::new())));
20static NAME_TYPES: LazyLock<RwLock<BTreeMap<String, (TypeId, &'static str)>>> =
21    LazyLock::new(Default::default);
22
23/// Sets the global namer.
24///
25/// All types will be named by this namer. Call this method before generating the OpenAPI schema.
26///
27/// # Example
28///
29/// ```rust
30/// # use salvo_oapi::extract::*;
31/// # use salvo_core::prelude::*;
32/// # #[tokio::main]
33/// # async fn main() {
34/// salvo_oapi::naming::set_namer(
35///     salvo_oapi::naming::FlexNamer::new()
36///         .short_mode(true)
37///         .generic_delimiter('_', '_'),
38/// );
39/// # }
40/// ```
41pub fn set_namer(namer: impl Namer) {
42    *GLOBAL_NAMER.write() = Box::new(namer);
43    NAME_TYPES.write().clear();
44}
45
46/// Reset global naming state to defaults.
47///
48/// This clears all registered type names and resets the namer to default `FlexNamer`.
49/// Primarily useful for testing to ensure test isolation.
50#[cfg(test)]
51pub fn reset_global_state() {
52    *GLOBAL_NAMER.write() = Box::new(FlexNamer::new());
53    NAME_TYPES.write().clear();
54}
55
56#[doc(hidden)]
57pub fn namer() -> RwLockReadGuard<'static, Box<dyn Namer>> {
58    GLOBAL_NAMER.read()
59}
60
61/// Get type info by name.
62pub fn type_info_by_name(name: &str) -> Option<(TypeId, &'static str)> {
63    NAME_TYPES.read().get(name).cloned()
64}
65
66/// Get registered name by rust type name (from `std::any::type_name`).
67///
68/// This searches through NAME_TYPES to find if a type with the given rust type name
69/// has been registered with a custom name.
70pub fn name_by_type_name(type_name: &str) -> Option<String> {
71    NAME_TYPES
72        .read()
73        .iter()
74        .find(|(_, (_, registered_type_name))| *registered_type_name == type_name)
75        .map(|(name, _)| name.clone())
76}
77
78/// Resolve generic type parameters to their registered names.
79///
80/// This function recursively processes a type name string and replaces any
81/// generic type parameters with their registered names from NAME_TYPES.
82///
83/// For example, if `CityDTO` is registered as `City`, then:
84/// - `Response<CityDTO>` becomes `Response<City>`
85/// - `Vec<HashMap<String, CityDTO>>` becomes `Vec<HashMap<String, City>>`
86#[must_use]
87pub fn resolve_generic_names(type_name: &str) -> String {
88    // First check if the entire type (without generics) has a registered name
89    if let Some(registered_name) = name_by_type_name(type_name) {
90        return registered_name;
91    }
92
93    // Find the position of the first '<' to separate base type from generic params
94    let Some(generic_start) = type_name.find('<') else {
95        // No generics, return as-is
96        return type_name.to_owned();
97    };
98
99    // Extract base type and generic part
100    let Some(base_type) = type_name.get(..generic_start) else {
101        return type_name.to_owned();
102    };
103    let Some(generic_part) = type_name.get(generic_start..) else {
104        return type_name.to_owned();
105    };
106
107    // Parse and resolve each generic parameter
108    let resolved_generic = resolve_generic_part(generic_part);
109
110    format!("{base_type}{resolved_generic}")
111}
112
113/// Parse generic part like `<A, B<C, D>, E>` and resolve each type parameter.
114fn resolve_generic_part(generic_part: &str) -> String {
115    if !generic_part.starts_with('<') || !generic_part.ends_with('>') {
116        return generic_part.to_owned();
117    }
118
119    // Remove outer < and >
120    let Some(inner) = generic_part
121        .strip_prefix('<')
122        .and_then(|generic_part| generic_part.strip_suffix('>'))
123    else {
124        return generic_part.to_owned();
125    };
126
127    // Split by top-level commas (not nested in <>)
128    let params = split_generic_params(inner);
129
130    let resolved_params: Vec<String> = params
131        .into_iter()
132        .map(|param| {
133            let param = param.trim();
134            // Check if this exact type has a registered name
135            if let Some(registered_name) = name_by_type_name(param) {
136                registered_name
137            } else if param.contains('<') {
138                // Recursively resolve nested generics
139                resolve_generic_names(param)
140            } else {
141                // Use short name for unregistered types (like primitive types)
142                // e.g., "alloc::string::String" -> "String"
143                short_type_name(param).to_owned()
144            }
145        })
146        .collect();
147
148    format!("<{}>", resolved_params.join(", "))
149}
150
151/// Split generic parameters at top-level commas, respecting nested angle brackets.
152fn split_generic_params(s: &str) -> Vec<&str> {
153    let mut result = Vec::new();
154    let mut depth = 0;
155    let mut start = 0;
156
157    for (i, c) in s.char_indices() {
158        match c {
159            '<' => depth += 1,
160            '>' => depth -= 1,
161            ',' if depth == 0 => {
162                if let Some(param) = s.get(start..i) {
163                    result.push(param);
164                }
165                start = i + 1;
166            }
167            _ => {}
168        }
169    }
170
171    // Don't forget the last segment
172    if start < s.len()
173        && let Some(param) = s.get(start..)
174    {
175        result.push(param);
176    }
177
178    result
179}
180
181/// Extract the short name from a fully qualified type path.
182///
183/// For example:
184/// - `alloc::string::String` -> `String`
185/// - `std::collections::HashMap` -> `HashMap`
186/// - `my_crate::module::MyType` -> `MyType`
187fn short_type_name(type_name: &str) -> &str {
188    // Find the last `::` and return everything after it
189    type_name
190        .rfind("::")
191        .and_then(|pos| type_name.get(pos + 2..))
192        .unwrap_or(type_name)
193}
194
195/// Set type info by name.
196pub fn set_name_type_info(
197    name: String,
198    type_id: TypeId,
199    type_name: &'static str,
200) -> Option<(TypeId, &'static str)> {
201    NAME_TYPES.write().insert(name, (type_id, type_name))
202}
203
204/// Assigns a name to a type and returns it.
205///
206/// If the type already has a registered name, the existing name is returned and
207/// `rule` is ignored. This function never panics: it consults the global namer
208/// and registers a fresh name on miss.
209pub fn assign_name<T: 'static>(rule: NameRule) -> String {
210    let type_id = TypeId::of::<T>();
211    let type_name = std::any::type_name::<T>();
212    for (name, (exist_id, _)) in NAME_TYPES.read().iter() {
213        if *exist_id == type_id {
214            return name.clone();
215        }
216    }
217    namer().assign_name(type_id, type_name, rule)
218}
219
220/// Get the name of the type. Panic if the name is not exist.
221pub fn get_name<T: 'static>() -> String {
222    let type_id = TypeId::of::<T>();
223    for (name, (exist_id, _)) in NAME_TYPES.read().iter() {
224        if *exist_id == type_id {
225            return name.clone();
226        }
227    }
228    panic!(
229        "Type not found in the name registry: {:?}",
230        std::any::type_name::<T>()
231    );
232}
233
234fn type_generic_part(type_name: &str) -> String {
235    if let Some(pos) = type_name.find('<') {
236        type_name.get(pos..).unwrap_or_default().to_owned()
237    } else {
238        String::new()
239    }
240}
241
242/// Resolve generic part and format it according to namer settings.
243fn resolve_and_format_generic_part(type_name: &str, short_mode: bool) -> String {
244    let generic_part = type_generic_part(type_name);
245    if generic_part.is_empty() {
246        return generic_part;
247    }
248
249    // Resolve registered names in generic parameters
250    let resolved = resolve_generic_part(&generic_part);
251
252    // Apply formatting (:: -> . for non-short mode, or strip module paths for short mode)
253    if short_mode {
254        let re = Regex::new(r"([^<>, ]*::)+").expect("Invalid regex");
255        re.replace_all(&resolved, "").into_owned()
256    } else {
257        resolved.replace("::", ".")
258    }
259}
260/// Namer is used to assign names to types.
261pub trait Namer: Sync + Send + 'static {
262    /// Assign name to type.
263    fn assign_name(&self, type_id: TypeId, type_name: &'static str, rule: NameRule) -> String;
264}
265
266/// A namer that generates wordy names.
267#[derive(Default, Clone, Debug)]
268pub struct FlexNamer {
269    short_mode: bool,
270    generic_delimiter: Option<(String, String)>,
271}
272impl FlexNamer {
273    /// Create a new FlexNamer.
274    #[must_use]
275    pub fn new() -> Self {
276        Default::default()
277    }
278
279    /// Set the short mode.
280    #[must_use]
281    pub fn short_mode(mut self, short_mode: bool) -> Self {
282        self.short_mode = short_mode;
283        self
284    }
285
286    /// Set the delimiter for generic types.
287    #[must_use]
288    pub fn generic_delimiter(mut self, open: impl Into<String>, close: impl Into<String>) -> Self {
289        self.generic_delimiter = Some((open.into(), close.into()));
290        self
291    }
292}
293impl Namer for FlexNamer {
294    fn assign_name(&self, type_id: TypeId, type_name: &'static str, rule: NameRule) -> String {
295        let name = match rule {
296            NameRule::Auto => {
297                // First resolve any registered names in generic parameters
298                let resolved_type_name = resolve_generic_names(type_name);
299
300                let mut base = if self.short_mode {
301                    let re = Regex::new(r"([^<>, ]*::)+").expect("Invalid regex");
302                    re.replace_all(&resolved_type_name, "").into_owned()
303                } else {
304                    resolved_type_name.replace("::", ".")
305                };
306                if let Some((open, close)) = &self.generic_delimiter {
307                    base = base.replace('<', open).replace('>', close);
308                }
309                let mut name = base.clone();
310                let mut count = 1;
311                while let Some(exist_id) = type_info_by_name(&name).map(|t| t.0) {
312                    if exist_id != type_id {
313                        count += 1;
314                        name = format!("{base}{count}");
315                    } else {
316                        break;
317                    }
318                }
319                name
320            }
321            NameRule::Force(force_name) => {
322                // Resolve registered names in generic parameters
323                let resolved_generic = resolve_and_format_generic_part(type_name, self.short_mode);
324
325                let mut base = if self.short_mode {
326                    // In short mode with Force, use the forced name + resolved generics
327                    format!("{force_name}{resolved_generic}")
328                } else {
329                    format!("{force_name}{resolved_generic}")
330                };
331                if let Some((open, close)) = &self.generic_delimiter {
332                    base = base.replace('<', open).replace('>', close);
333                }
334                let mut name = base.clone();
335                let mut count = 1;
336                while let Some((exist_id, exist_name)) = type_info_by_name(&name) {
337                    if exist_id != type_id {
338                        count += 1;
339                        tracing::error!("Duplicate name for types: {}, {}", exist_name, type_name);
340                        name = format!("{base}{count}");
341                    } else {
342                        break;
343                    }
344                }
345                name
346            }
347        };
348        set_name_type_info(name.clone(), type_id, type_name);
349        name
350    }
351}
352
353#[cfg(test)]
354mod tests {
355    use serial_test::serial;
356
357    #[test]
358    #[serial]
359    fn test_name() {
360        use super::*;
361
362        // Reset global state to ensure deterministic test results
363        reset_global_state();
364
365        struct MyString;
366        mod nest {
367            pub(crate) struct MyString;
368        }
369
370        let name = assign_name::<String>(NameRule::Auto);
371        assert_eq!(name, "alloc.string.String");
372        let name = assign_name::<Vec<String>>(NameRule::Auto);
373        assert_eq!(name, "alloc.vec.Vec<alloc.string.String>");
374
375        let name = assign_name::<MyString>(NameRule::Auto);
376        assert!(
377            name.contains("MyString") && !name.contains("nest"),
378            "Expected name containing 'MyString' but not 'nest', got: {name}"
379        );
380        let name = assign_name::<nest::MyString>(NameRule::Auto);
381        assert!(
382            name.contains("nest") && name.contains("MyString"),
383            "Expected name containing 'nest.MyString', got: {name}"
384        );
385    }
386
387    #[test]
388    #[serial]
389    fn test_resolve_generic_names() {
390        use super::*;
391
392        // Reset global state to ensure deterministic test results
393        reset_global_state();
394
395        // Simulate registering CityDTO as "City"
396        let city_type_name = "test_module::CityDTO";
397        set_name_type_info(
398            "City".to_owned(),
399            TypeId::of::<()>(), // dummy TypeId
400            city_type_name,
401        );
402
403        // Test resolve_generic_names with registered type
404        let resolved = resolve_generic_names("Response<test_module::CityDTO>");
405        assert_eq!(resolved, "Response<City>");
406
407        // Test nested generics - unregistered types get short names
408        let resolved = resolve_generic_names("Vec<HashMap<String, test_module::CityDTO>>");
409        assert_eq!(resolved, "Vec<HashMap<String, City>>");
410
411        // Test multiple generic parameters
412        let resolved = resolve_generic_names("Tuple<test_module::CityDTO, test_module::CityDTO>");
413        assert_eq!(resolved, "Tuple<City, City>");
414    }
415
416    #[test]
417    #[serial]
418    fn test_resolve_primitive_types() {
419        use super::*;
420
421        // Reset global state to ensure deterministic test results
422        reset_global_state();
423
424        // Test with primitive types (not registered, should use short names in generic params)
425        let resolved = resolve_generic_names("Response<alloc::string::String>");
426        assert_eq!(resolved, "Response<String>");
427
428        // Note: The base type (Vec) keeps its path, only generic params are shortened
429        // FlexNamer::assign_name handles the full path transformation later
430        let resolved = resolve_generic_names("Vec<alloc::vec::Vec<alloc::string::String>>");
431        assert_eq!(resolved, "Vec<alloc::vec::Vec<String>>");
432
433        // Test HashMap with primitive types
434        let resolved =
435            resolve_generic_names("std::collections::HashMap<alloc::string::String, i32>");
436        assert_eq!(resolved, "std::collections::HashMap<String, i32>");
437
438        // Test that nested generic base types are also shortened in their generics
439        let resolved = resolve_generic_names("Option<Vec<alloc::string::String>>");
440        assert_eq!(resolved, "Option<Vec<String>>");
441    }
442
443    #[test]
444    fn test_short_type_name() {
445        use super::*;
446
447        assert_eq!(short_type_name("alloc::string::String"), "String");
448        assert_eq!(short_type_name("std::collections::HashMap"), "HashMap");
449        assert_eq!(short_type_name("MyType"), "MyType");
450        assert_eq!(short_type_name("my_crate::module::submodule::Type"), "Type");
451    }
452
453    #[test]
454    fn test_split_generic_params() {
455        use super::*;
456
457        let params = split_generic_params("A, B, C");
458        assert_eq!(params, vec!["A", " B", " C"]);
459
460        let params = split_generic_params("A<X, Y>, B, C<Z>");
461        assert_eq!(params, vec!["A<X, Y>", " B", " C<Z>"]);
462
463        let params = split_generic_params("A<X<Y, Z>>, B");
464        assert_eq!(params, vec!["A<X<Y, Z>>", " B"]);
465    }
466
467    #[test]
468    #[serial]
469    fn test_assign_name_with_generic_resolution() {
470        use super::*;
471
472        // Reset global state to ensure deterministic test results
473        reset_global_state();
474
475        // Define unique test types for this test to avoid conflicts with other tests
476        mod test_generic_resolution {
477            pub(super) struct CityDTO;
478            pub(super) struct Response<T>(std::marker::PhantomData<T>);
479            pub(super) struct Wrapper<T>(std::marker::PhantomData<T>);
480        }
481        use test_generic_resolution::*;
482
483        // First, register CityDTO with a custom name "City"
484        let city_name = assign_name::<CityDTO>(NameRule::Force("City"));
485        assert_eq!(city_name, "City");
486
487        // Now register Response<CityDTO> with Force("Response")
488        // It should resolve CityDTO to "City" in the generic parameter
489        let response_name = assign_name::<Response<CityDTO>>(NameRule::Force("Response"));
490        assert_eq!(response_name, "Response<City>");
491
492        // Test with Auto mode - should also resolve generic parameters
493        let wrapper_name = assign_name::<Wrapper<CityDTO>>(NameRule::Auto);
494        // The base type will have full path, but CityDTO should be resolved to City
495        assert!(
496            wrapper_name.contains("<City>"),
497            "Expected wrapper name to contain '<City>', got: {wrapper_name}"
498        );
499    }
500
501    #[test]
502    #[serial]
503    fn test_assign_name_with_primitive_generics() {
504        use super::*;
505
506        // Reset global state to ensure deterministic test results
507        reset_global_state();
508
509        mod test_primitive_generics {
510            pub(super) struct Response<T>(std::marker::PhantomData<T>);
511        }
512        use test_primitive_generics::*;
513
514        // Test Response<String> with Force("Response")
515        // String is not registered, but should be shortened to "String"
516        let response_name = assign_name::<Response<String>>(NameRule::Force("Response"));
517        assert_eq!(response_name, "Response<String>");
518
519        // Test Response<Vec<String>> - nested generics with primitives
520        let response_vec_name =
521            assign_name::<Response<Vec<String>>>(NameRule::Force("ResponseVec"));
522        assert!(
523            response_vec_name.contains("<String>"),
524            "Expected name to contain '<String>', got: {response_vec_name}"
525        );
526    }
527}