Skip to main content

fast_router/
router.rs

1use hyper::server::conn::AddrStream;
2use hyper::service::{make_service_fn, service_fn};
3use hyper::{Body, Request, Response, Server};
4use log::{debug, error, info, trace, warn};
5use std::collections::HashMap;
6use std::fmt::Debug;
7use std::str::FromStr;
8use std::sync::Arc;
9use std::{convert::Infallible, net::SocketAddr};
10use urlencoding::decode;
11
12use regex::Regex;
13
14/// 请求处理函数
15type Handler = fn(Context);
16
17#[derive(Debug, PartialEq, Eq)]
18pub enum Method {
19    TRACE,
20    HEAD,
21    GET,
22    POST,
23    PUT,
24    PATCH,
25    DELETE,
26    OPTIONS,
27    ANY,
28}
29
30impl From<Method> for String {
31    fn from(v: Method) -> Self {
32        format!("{:?}", v)
33    }
34}
35
36impl From<&str> for Method {
37    fn from(s: &str) -> Self {
38        match s.to_uppercase().as_str() {
39            "TRACE" => Self::TRACE,
40            "HEAD" => Self::HEAD,
41            "GET" => Self::GET,
42            "POST" => Self::POST,
43            "PUT" => Self::PUT,
44            "PATCH" => Self::PATCH,
45            "DELETE" => Self::DELETE,
46            "OPTIONS" => Self::OPTIONS,
47            "ANY" => Self::ANY,
48            v => panic!("不存在这个类型:{}", v),
49        }
50    }
51}
52
53/// 包含了请求的所有信息以及用户自定义信息
54#[derive(Default)]
55pub struct Context {
56    params: HashMap<String, String>,
57    queries: HashMap<String, String>,
58}
59
60impl Context {
61    /// 获取路由规则中的命名参数值
62    pub fn param<T>(&self, name: &str) -> Option<T>
63    where
64        T: FromStr,
65        <T as FromStr>::Err: Debug,
66    {
67        self.params
68            .get(name.into())
69            .map(|v| v.as_str().parse().unwrap())
70    }
71}
72
73#[derive(Default, Debug)]
74pub struct Router {
75    // 请求处理之前的过滤器
76    before_fileters: Vec<Route>,
77
78    // 处理请求的函数
79    routes: Vec<Route>,
80
81    // 处理之后的过滤器
82    after_fileters: Vec<Route>,
83
84    // 是否忽略路径后面的斜线,默认忽略
85    has_slash: bool,
86}
87
88impl Router {
89    async fn shutdown_signal() {
90        tokio::signal::ctrl_c()
91            .await
92            .expect("安装 CTRL+C 处理器失败");
93    }
94
95    /// 静态资源目录
96    pub fn static_dir(path: &str, dir: &str) {}
97
98    /// 静态文件
99    pub fn static_file(uri: &str, filepath: &str) {}
100
101    async fn handle(
102        router: Arc<Router>,
103        addr: SocketAddr,
104        req: Request<Body>,
105    ) -> Result<Response<Body>, Infallible> {
106        let mut context = Context::default();
107
108        //
109        let routes = router.match_route(req.method().as_str().into(), req.uri().path());
110        trace!("匹配到的路由:{:?}", routes);
111
112        // 如果匹配成功,从路由中提取参数的值,保存到context中
113        if let Some(r) = routes {
114            if let Some(r) = r.route {
115                trace!("地址:{}", req.uri().path());
116                trace!("路由规则:{:?}", r);
117                if let Some(c) = r.re.captures(req.uri().path()) {
118                    for r in &r.param_names {
119                        match decode(c.name(r.as_str()).unwrap().as_str()) {
120                            Ok(value) => {
121                                context.params.insert(r.into(), value.to_string());
122                            }
123                            Err(e) => {
124                                warn!("路由参数值urldecode解码出错:{:?}", e)
125                            }
126                        }
127                    }
128                }
129
130                (r.handler)(context);
131            }
132        }
133
134        Ok(Response::new(Body::from("Hello World")))
135    }
136
137    /// 启动服务器
138    pub fn run(self, host: &str) {
139        let rt = tokio::runtime::Runtime::new().unwrap();
140        rt.block_on(async {
141            let addr = match SocketAddr::from_str(host) {
142                Ok(v) => v,
143                Err(_) => {
144                    error!("解析地址失败,请确认格式为(ip:port),你的地址是:{}", host);
145                    return;
146                }
147            };
148
149            // 路由表信息
150            let router: Arc<Router> = Arc::new(self);
151
152            debug!("路由表:{:#?}", router.clone());
153
154            // 创建service处理每个请求
155            let make_service = make_service_fn(move |conn: &AddrStream| {
156                // 客户端地址信息
157                let addr = conn.remote_addr();
158
159                // 每个请求都克隆一个路由信息,传入给处理函数
160                let router = router.clone();
161                let service = service_fn(move |req| Self::handle(router.clone(), addr, req));
162                async move { Ok::<_, Infallible>(service) }
163            });
164            let server = Server::bind(&addr).serve(make_service);
165
166            let graceful = server.with_graceful_shutdown(Self::shutdown_signal());
167
168            info!("启动成功: {}", host);
169
170            if let Err(e) = graceful.await {
171                eprintln!("server error: {}", e);
172            }
173        });
174    }
175
176    /// 启动tls服务器
177    pub fn run_tls(host: &str, pem: &str, key: &str) {}
178}
179
180impl Router {
181    pub fn new() -> Self {
182        Self::default()
183    }
184
185    pub fn group(&mut self, path: &str) -> RouterGroup {
186        RouterGroup::new(path, self)
187    }
188    fn add(&mut self, method: Method, path: &str, handler: Handler) {
189        // assert!(path.len() > 0);
190
191        let route = Route::new(method, path.to_owned(), handler, None, self.has_slash);
192
193        self.routes.push(route);
194    }
195
196    /// 区分路径最后的斜线,默认不区分
197    pub fn has_slash(&mut self) {
198        self.has_slash = true;
199    }
200
201    /// 根据用户请求地址,匹配路由
202    fn match_route(&self, method: Method, path: &str) -> Option<MatchedRoute> {
203        // 如果忽略地址最后的斜线,并且长度不是0,就去掉后面的斜线
204        let path = if !self.has_slash && path.len() > 0 && &path[path.len() - 1..] == "/" {
205            &path[..path.len() - 1]
206        } else {
207            path
208        };
209
210        let mut params: HashMap<String, String> = HashMap::new();
211
212        trace!("查找路由:{}", path);
213        // 寻找匹配的路由
214        let mut matched_route = None;
215        for route in &self.routes {
216            if (method == route.method || route.method == Method::ANY) && route.re.is_match(path) {
217                // 找到匹配的路由
218                matched_route = Some(route);
219
220                // 提取参数
221                let cps = route.re.captures(path).unwrap();
222                trace!("参数列表:{:?}", route.param_names);
223                for name in &route.param_names {
224                    params.insert(
225                        name.to_string(),
226                        cps.name(name.as_str()).unwrap().as_str().into(),
227                    );
228                }
229
230                break;
231            }
232        }
233
234        // 寻找前置过滤器
235        let mut before_filters = vec![];
236        for route in &self.before_fileters {
237            if (method == route.method || route.method == Method::ANY) && route.re.is_match(path) {
238                // 找到匹配的过滤器,可能会有多个
239                before_filters.push(route);
240
241                // 提取参数
242                let cps = route.re.captures(path).unwrap();
243                trace!("参数列表:{:?}", route.param_names);
244                for name in &route.param_names {
245                    params.insert(
246                        name.to_string(),
247                        cps.name(name.as_str()).unwrap().as_str().into(),
248                    );
249                }
250            }
251        }
252
253        // 输出所有参数值
254        trace!("路径中的参数值:{:?}", params);
255
256        // 寻找后置过滤器
257        let mut after_filters = vec![];
258        for route in &self.after_fileters {
259            if (method == route.method || route.method == Method::ANY) && route.re.is_match(path) {
260                // 找到匹配的过滤器,可能会有多个
261                after_filters.push(route);
262            }
263        }
264
265        // 如果都没找到
266        if matched_route.is_none() && before_filters.len() == 0 && after_filters.len() == 0 {
267            return None;
268        }
269
270        Some(MatchedRoute {
271            before: before_filters,
272            route: matched_route,
273            after: after_filters,
274            params: None,
275        })
276    }
277
278    /// 添加前置中间件
279    pub fn before(&mut self, method: Method, path: &str, handler: Handler) {
280        let route = Route::new_filter(method, path.to_owned(), handler, None, self.has_slash);
281
282        self.before_fileters.push(route);
283    }
284
285    /// 添加后置中间件
286    pub fn after(&mut self, method: Method, path: &str, handler: Handler) {
287        let route = Route::new_filter(method, path.to_owned(), handler, None, self.has_slash);
288
289        self.after_fileters.push(route);
290    }
291
292    /// 封装各类请求
293    pub fn get(&mut self, path: &str, handler: Handler) {
294        self.add(Method::GET, path, handler);
295    }
296
297    pub fn post(&mut self, path: &str, handler: Handler) {
298        self.add(Method::POST, path, handler);
299    }
300
301    pub fn trace(&mut self, path: &str, handler: Handler) {
302        self.add(Method::TRACE, path, handler);
303    }
304
305    pub fn head(&mut self, path: &str, handler: Handler) {
306        self.add(Method::HEAD, path, handler);
307    }
308
309    pub fn put(&mut self, path: &str, handler: Handler) {
310        self.add(Method::PUT, path, handler);
311    }
312
313    pub fn patch(&mut self, path: &str, handler: Handler) {
314        self.add(Method::PATCH, path, handler);
315    }
316
317    pub fn delete(&mut self, path: &str, handler: Handler) {
318        self.add(Method::DELETE, path, handler);
319    }
320
321    pub fn options(&mut self, path: &str, handler: Handler) {
322        self.add(Method::OPTIONS, path, handler);
323    }
324
325    pub fn any(&mut self, path: &str, handler: Handler) {
326        self.add(Method::ANY, path, handler);
327    }
328}
329
330/// 根据地址匹配到的路由
331#[derive(Debug)]
332struct MatchedRoute<'a> {
333    before: Vec<&'a Route>,
334    route: Option<&'a Route>,
335    after: Vec<&'a Route>,
336    // 地址参数
337    params: Option<HashMap<String, String>>,
338}
339
340pub struct RouterGroup<'a> {
341    path: String,
342    router: &'a mut Router,
343}
344
345impl<'a> RouterGroup<'a> {
346    fn new(path: &str, router: &'a mut Router) -> Self {
347        Self {
348            path: path.to_string(),
349            router: router,
350        }
351    }
352
353    /// 处理两个地址相加,防止地址相加出现两个斜线或没有斜线
354    fn concat_path(path1: &str, path2: &str) -> String {
355        // 两个路径除去斜杠之后的长度
356        let l1 = path1.replace("/", "").len();
357        let l2 = path2.replace("/", "").len();
358
359        if l1 == 0 && l2 > 0 {
360            // 如果前面的地址为空,返回第二个地址
361            return path2.to_string();
362        } else if l2 == 0 && l1 > 0 {
363            // 如果后面的地址为空,返回第一个
364            return path1.to_string();
365        } else if l1 == 0 && l2 == 0 {
366            return "".to_string();
367        }
368
369        match (&path1[path1.len() - 1..], &path2[0..1]) {
370            // 两个斜线就去掉一个
371            ("/", "/") => path1.to_string() + &path2[1..],
372
373            // 没有斜线就添加一个
374            (p1, p2) if p1 != "/" && p2 != "/" => path1.to_string() + "/" + path2,
375
376            // 一个斜线就直接连起来
377            _ => path1.to_string() + path2,
378        }
379    }
380
381    fn add(&mut self, method: Method, path: &str, handler: Handler) {
382        let path = Self::concat_path(self.path.as_str(), path);
383
384        self.router.add(method, path.as_str(), handler);
385    }
386
387    /// 生成一个分组
388    pub fn group(&mut self, path: &str) -> RouterGroup {
389        let path = Self::concat_path(self.path.as_str(), path);
390        RouterGroup::new(path.as_str(), self.router)
391    }
392
393    /// 添加前置处理器
394    pub fn before(&mut self, method: Method, path: &str, handler: Handler) {
395        let path = Self::concat_path(self.path.as_str(), path);
396
397        self.router.before(method, path.as_str(), handler);
398    }
399
400    /// 添加后置处理器
401    pub fn after(&mut self, method: Method, path: &str, handler: Handler) {
402        let path = Self::concat_path(self.path.as_str(), path);
403
404        self.router.after(method, path.as_str(), handler);
405    }
406
407    /// 封装各类请求
408    pub fn get(&mut self, path: &str, handler: Handler) {
409        self.add(Method::GET, path, handler);
410    }
411
412    pub fn post(&mut self, path: &str, handler: Handler) {
413        self.add(Method::POST, path, handler);
414    }
415
416    pub fn trace(&mut self, path: &str, handler: Handler) {
417        self.add(Method::TRACE, path, handler);
418    }
419
420    pub fn head(&mut self, path: &str, handler: Handler) {
421        self.add(Method::HEAD, path, handler);
422    }
423
424    pub fn put(&mut self, path: &str, handler: Handler) {
425        self.add(Method::PUT, path, handler);
426    }
427
428    pub fn patch(&mut self, path: &str, handler: Handler) {
429        self.add(Method::PATCH, path, handler);
430    }
431
432    pub fn delete(&mut self, path: &str, handler: Handler) {
433        self.add(Method::DELETE, path, handler);
434    }
435
436    pub fn options(&mut self, path: &str, handler: Handler) {
437        self.add(Method::OPTIONS, path, handler);
438    }
439
440    pub fn any(&mut self, path: &str, handler: Handler) {
441        self.add(Method::ANY, path, handler);
442    }
443}
444
445/// 路径的一个路由信息
446#[derive(Debug)]
447struct Route {
448    // 请求方式,不区分大小写
449    method: Method,
450
451    // 路由名字,用于生成url
452    name: Option<String>,
453
454    // 路由规则
455    path: String,
456
457    // 路径编译之后的正则表达式对象
458    re: Regex,
459
460    // 处理函数
461    handler: Handler,
462
463    // 是否保留路径最后的斜线
464    has_slash: bool,
465
466    // 路径中的分组命名
467    param_names: Vec<String>,
468}
469
470impl Route {
471    /// 创建路由
472    fn new(
473        method: Method,
474        path: String,
475        handler: Handler,
476        name: Option<String>,
477        has_slash: bool,
478    ) -> Self {
479        Self::build(method, path, handler, name, has_slash, false)
480    }
481
482    /// 创建过滤器
483    fn new_filter(
484        method: Method,
485        path: String,
486        handler: Handler,
487        name: Option<String>,
488        has_slash: bool,
489    ) -> Self {
490        Self::build(method, path, handler, name, has_slash, true)
491    }
492    fn build(
493        method: Method,
494        path: String,
495        handler: Handler,
496        name: Option<String>,
497        has_slash: bool,
498
499        // 是否是过滤器,如果是过滤器,正则结尾不需要$,这样只要匹配一部分路径即可
500        is_filter: bool,
501    ) -> Self {
502        // 如果忽不保留路径最后的斜线,就去掉,否则就不变
503        let path = if !has_slash && !path.is_empty() && &path[path.len() - 1..] == "/" {
504            path[..path.len() - 1].to_string()
505        } else {
506            path
507        };
508
509        // 把自定义类型的路由转换成 正则路由
510        let path = Self::path_param_type_to_regex(path.as_str());
511
512        // 把正则路由转换成 在组名的正则路由
513        let path_and_names = Self::path2regex(path.as_str());
514
515        trace!("路由参数列表:{:?}", path_and_names.1);
516
517        let mut re_str = format!("^{}", path_and_names.0);
518
519        // 如果不是过滤器,需要匹配整个地址
520        if !is_filter {
521            re_str += "$";
522        }
523
524        let re = Regex::new(re_str.as_str()).unwrap();
525        Self {
526            method,
527            path: path.clone(),
528            name,
529            re,
530            handler,
531            has_slash,
532            param_names: path_and_names.1,
533        }
534    }
535
536    /// 返回对应类型的正则
537    #[inline]
538    fn type_to_regex(type_name: &str) -> String {
539        // i8|u8|i32|u32|i64|u64|i128|u128|bool
540        match type_name {
541            "i8" => r"[\-]{0,1}\d{1,3}",
542            "i16" => r"[\-]{0,1}\d{1,5}",
543            "i32" => r"[\-]{0,1}\d{1,10}",
544            "i64" => r"[\-]{0,1}\d{1,19}",
545            "i128" => r"[\-]{0,1}\d{1,39}",
546            "u8" => r"\d{1,3}",
547            "u16" => r"\d{1,5}",
548            "u32" => r"\d{1,10}",
549            "u64" => r"\d{1,20}",
550            "u128" => r"\d{1,39}",
551            "bool" => r"true|false",
552            v => {
553                panic!("路由不支持该参数类型:{}", v);
554            }
555        }
556        .to_string()
557    }
558
559    /// 把路由规则转换成正则表达式格式
560    ///
561    /// 例如:  /user/:id:usize/:page:usize
562    /// 转换成:/user/:id:(\d+)/:page:(\d+)
563    ///
564    /// 返回值说明:
565    /// 返回的第一个值是转换后的正则路由,第二个参数是正则路由中的命名参数名字
566    /// 比如:/:user/:id  第二个参数就返回["user","id"]
567    ///
568    #[inline]
569    fn path_param_type_to_regex(path: &str) -> String {
570        let mut p = String::new();
571
572        let re = Regex::new(r#"^:(?P<name>[a-zA-a_]{1}[a-zA-Z_0-9]*?):(?P<type>i32|u32|i8|u8|i64|u64|i128|u128|bool)$"#).unwrap();
573
574        for node in path.split("/") {
575            if node.is_empty() {
576                continue;
577            }
578
579            if re.is_match(node) {
580                let cms = re.captures(node).unwrap();
581                let name = cms.name("name").unwrap().as_str();
582                let tp = cms.name("type").unwrap().as_str();
583
584                let type_reg = Self::type_to_regex(tp);
585                p += format!("/:{}:({})", name, type_reg).as_str();
586            } else if &node[0..1] == ":" && &node[node.len() - 1..] != ")" {
587                // 如果不匹配,以 : 开头,表示没写类型,默认匹配任意字符串
588                // 也就是这种情况 /:name/
589                p = p + "/" + node + r#":([\w\-%_\.~:;'"@=+,]+)"#;
590            } else {
591                // 自定义正则就不改变
592                p = p + "/" + node;
593            }
594        }
595        // 最后如果有 / 也要加上
596        if &path[path.len() - 1..] == "/" {
597            p += "/";
598        }
599
600        p
601    }
602
603    /// 把正则路由转换成,命名组正则表达式,如果是自定义类型的,需要先调用 path_param_type_to_regex 函数来处理成正则路由
604    ///
605    /// 例如把 /admin/:name:([^/]+)/:id:(\d+)
606    /// 转换成 /admin/(?P<name>[^/]+)/(?P<id>\d+)
607    ///
608    #[inline]
609    fn path2regex(path: &str) -> (String, Vec<String>) {
610        let mut p = String::new();
611
612        let re = Regex::new(r#"^:(?P<name>[a-zA-a_]{1}[a-zA-Z_0-9]*?):\((?P<reg>.*)\)$"#).unwrap();
613
614        let mut names = vec![];
615
616        for node in path.split("/") {
617            if node.is_empty() {
618                continue;
619            }
620
621            if re.is_match(node) {
622                let cms = re.captures(node).unwrap();
623                let name = cms.name("name").unwrap().as_str();
624                names.push(name.to_string());
625
626                p += re
627                    .replace(node, "/(?P<${name}>${reg})")
628                    .to_string()
629                    .as_str();
630            } else {
631                p += "/";
632                p += node;
633            }
634        }
635        // 最后如果有 / 也要加上
636        if &path[path.len() - 1..] == "/" {
637            p += "/";
638        }
639
640        (p, names)
641    }
642}
643
644#[cfg(test)]
645mod tests {
646
647    use regex::Regex;
648
649    use crate::router::Route;
650
651    use super::Method;
652    use super::Router;
653    use super::RouterGroup;
654
655    #[test]
656    fn test_concat() {
657        assert_eq!("a/b".to_string(), RouterGroup::concat_path("a/", "/b"));
658        assert_eq!("a/b".to_string(), RouterGroup::concat_path("a", "/b"));
659        assert_eq!("a/b".to_string(), RouterGroup::concat_path("a/", "b"));
660        assert_eq!("a/b".to_string(), RouterGroup::concat_path("a", "b"));
661        assert_eq!("a".to_string(), RouterGroup::concat_path("a", ""));
662        assert_eq!("".to_string(), RouterGroup::concat_path("", ""));
663        assert_eq!("b".to_string(), RouterGroup::concat_path("", "b"));
664    }
665
666    #[test]
667    fn test_router() {
668        let mut r = Router::default();
669        r.has_slash();
670
671        r.post("/:user/:id", |c| {});
672
673        let mut g = r.group("/:v1");
674        {
675            g.get("admin/:name:i32", |mut c| {});
676            g.post("/admin/u32/", |c| {});
677
678            let mut g1 = g.group("/test1");
679            {
680                g1.put("admin/login/", |c| {});
681                g1.delete("/admin/login1/", |c| {});
682            }
683        }
684
685        let mut g = r.group("/a1/");
686        {
687            g.add(Method::GET, "/admin/login", |c| {});
688            g.add(Method::DELETE, "/admin/login1", |c| {});
689
690            let mut g1 = g.group("/test1/");
691            {
692                g1.add(Method::OPTIONS, "/admin/login", |c| {});
693                g1.add(Method::ANY, "/admin/login1", |c| {});
694            }
695        }
696
697        for v in r.routes {
698            println!("{:?}", v);
699        }
700    }
701
702    #[test]
703    fn test_regex() {
704        let re = Regex::new(r"^/user/(?P<name>\w+)/(?P<id>\d{1,10})$").unwrap();
705
706        let v = re.captures("/user/zhangsan/123").unwrap();
707        let r = v.name("name").unwrap();
708        println!("{:?}", r.as_str());
709        let r = v.name("id").unwrap();
710        println!("{:?}", r.as_str());
711    }
712
713    #[test]
714    fn test_path2regex() {
715        /*
716        把     /admin/:name:([^/]+)/:id:(\d+)
717            转换成 /admin/(?P<name>[^/]+)/(?P<id>\d+)
718        */
719
720        let s = r#"/admin/:name:(.*+?)/info/:id:(\d+?)/name/"#;
721        let p = Route::path2regex(s);
722
723        assert_eq!(
724            r"/admin/(?P<name>.*+?)/info/(?P<id>\d+?)/name/",
725            p.0.as_str()
726        );
727    }
728
729    #[test]
730    fn test_path_param_to_regex() {
731        let path = "/user/:id:(.*)/:page:u32";
732        let p = Route::path_param_type_to_regex(path);
733        assert_eq!(r"/user/:id:(.*)/:page:(\d{1,10})", p.as_str());
734    }
735
736    #[test]
737    fn test_router_match() {
738        let mut r = Router::default();
739
740        let mut g = r.group("/v1");
741        {
742            g.before(Method::ANY, "admin/:name", |c| {});
743            g.after(Method::ANY, "admin/:name", |c| {});
744
745            g.any("admin/:name:i32", |c| {});
746            g.any("admin/:name", |c| {});
747            g.any("admin/:name/:id:u32", |c| {});
748            g.post("/admin/login1/", |c| {});
749        }
750
751        let route = r.match_route(Method::GET, "/v1/admin/zhang山/23423");
752
753        // (route.as_ref().unwrap().route.unwrap().handler)("xxx".to_string());
754    }
755
756    #[test]
757    fn test_filters() {}
758}