1use core::str::Utf8Error;
66use percent_encoding::percent_decode;
67use std::{collections::BTreeMap, error::Error, fmt, str::Chars};
68
69#[derive(Debug)]
71pub enum ParseError {
72 InvalidDriver,
74 InvalidParams,
76 InvalidPath,
78 InvalidPort,
80 InvalidProtocol,
82 InvalidSocket,
84 MissingAddress,
86 MissingHost,
88 MissingProtocol,
90 MissingSocket,
92 Utf8Error(Utf8Error),
94}
95
96impl From<Utf8Error> for ParseError {
97 fn from(err: Utf8Error) -> Self {
98 Self::Utf8Error(err)
99 }
100}
101
102impl fmt::Display for ParseError {
103 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
104 match *self {
105 Self::InvalidDriver => write!(f, "invalid driver"),
106 Self::InvalidParams => write!(f, "invalid params"),
107 Self::InvalidPath => write!(f, "invalid absolute path"),
108 Self::InvalidPort => write!(f, "invalid port number"),
109 Self::InvalidProtocol => write!(f, "invalid protocol"),
110 Self::InvalidSocket => write!(f, "invalid socket"),
111 Self::MissingAddress => write!(f, "missing address"),
112 Self::MissingHost => write!(f, "missing host"),
113 Self::MissingProtocol => write!(f, "missing protocol"),
114 Self::MissingSocket => write!(f, "missing unix domain socket"),
115 Self::Utf8Error(ref err) => write!(f, "UTF-8 error: {err}"),
116 }
117 }
118}
119
120impl Error for ParseError {}
121
122impl fmt::Display for DSN {
123 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
124 use percent_encoding::{NON_ALPHANUMERIC, utf8_percent_encode};
125
126 write!(f, "{}://", self.driver)?;
127
128 if let Some(ref username) = self.username {
130 let encoded_user = utf8_percent_encode(username, NON_ALPHANUMERIC);
131 write!(f, "{encoded_user}")?;
132
133 if let Some(ref password) = self.password {
134 let encoded_pass = utf8_percent_encode(password, NON_ALPHANUMERIC);
135 write!(f, ":{encoded_pass}")?;
136 }
137 write!(f, "@")?;
138 }
139
140 write!(f, "{}({})", self.protocol, self.address)?;
142
143 if let Some(ref database) = self.database {
145 write!(f, "/{database}")?;
146 }
147
148 if !self.params.is_empty() {
150 write!(f, "?")?;
151 let params: Vec<String> = self
152 .params
153 .iter()
154 .map(|(k, v)| format!("{k}={v}"))
155 .collect();
156 write!(f, "{}", params.join("&"))?;
157 }
158
159 Ok(())
160 }
161}
162
163#[derive(Clone, Debug, Default, PartialEq, Eq, Hash)]
178pub struct DSN {
179 pub driver: String,
181 pub username: Option<String>,
183 pub password: Option<String>,
185 pub protocol: String,
187 pub address: String,
189 pub host: Option<String>,
191 pub port: Option<u16>,
193 pub database: Option<String>,
195 pub socket: Option<String>,
197 pub params: BTreeMap<String, String>,
199}
200
201pub fn parse(input: &str) -> Result<DSN, ParseError> {
258 let mut dsn = DSN::default();
260
261 let chars = &mut input.chars();
263
264 dsn.driver = get_driver(chars)?;
266
267 let (user, pass) = get_username_password(chars)?;
269 if !user.is_empty() {
270 dsn.username = Some(user);
271 }
272 if !pass.is_empty() {
273 dsn.password = Some(pass);
274 }
275
276 dsn.protocol = get_protocol(chars)?;
278
279 dsn.address = get_address(chars)?;
281
282 match dsn.protocol.as_str() {
283 "unix" => {
284 if !dsn.address.starts_with('/') {
285 return Err(ParseError::InvalidSocket);
286 }
287 dsn.socket = Some(dsn.address.clone());
288 }
289 "file" => {
290 if !dsn.address.starts_with('/') {
291 return Err(ParseError::InvalidPath);
292 }
293 }
294 _ => {
295 let (host, port) = get_host_port(&dsn.address)?;
296 dsn.host = Some(host);
297
298 if !port.is_empty() {
299 dsn.port = Some(port.parse::<u16>().map_err(|_| ParseError::InvalidPort)?);
300 }
301 }
302 }
303
304 let database = get_database(chars);
306 if !database.is_empty() {
307 dsn.database = Some(database);
308 }
309
310 let params = chars.as_str();
311 if !params.is_empty() {
312 dsn.params = get_params(chars.as_str())?;
313 }
314
315 Ok(dsn)
316}
317
318fn get_driver(chars: &mut Chars) -> Result<String, ParseError> {
327 let mut driver = String::new();
328 while let Some(c) = chars.next() {
329 if c == ':' {
330 if chars.next() == Some('/') && chars.next() == Some('/') {
331 break;
332 }
333 return Err(ParseError::InvalidDriver);
334 }
335 driver.push(c);
336 }
337 Ok(driver)
338}
339
340fn get_username_password(chars: &mut Chars) -> Result<(String, String), ParseError> {
350 let mut username = String::new();
351 let mut password = String::new();
352 let mut has_password = true;
353
354 for c in chars.by_ref() {
356 match c {
357 '@' => {
358 has_password = false;
359 break;
360 }
361 ':' => {
362 break;
363 }
364 _ => username.push(c),
365 }
366 }
367
368 username = percent_decode(username.as_bytes()).decode_utf8()?.into();
369
370 if has_password {
372 for c in chars {
373 match c {
374 '@' => break,
375 _ => password.push(c),
376 }
377 }
378 password = percent_decode(password.as_bytes()).decode_utf8()?.into();
379 }
380
381 Ok((username, password))
382}
383
384fn get_protocol(chars: &mut Chars) -> Result<String, ParseError> {
393 let mut protocol = String::new();
394 for c in chars {
395 match c {
396 '(' => {
397 if protocol.is_empty() {
398 return Err(ParseError::MissingProtocol);
399 }
400 break;
401 }
402 _ => protocol.push(c),
403 }
404 }
405 Ok(protocol)
406}
407
408fn get_address(chars: &mut Chars) -> Result<String, ParseError> {
417 let mut address = String::new();
418 for c in chars {
419 match c {
420 ')' => {
421 if address.is_empty() {
422 return Err(ParseError::MissingAddress);
423 }
424 break;
425 }
426 _ => address.push(c),
427 }
428 }
429 Ok(address)
430}
431
432fn get_host_port(address: &str) -> Result<(String, String), ParseError> {
442 let mut host = String::new();
443 let mut chars = address.chars();
444
445 for c in chars.by_ref() {
447 match c {
448 ':' => {
449 if host.is_empty() {
450 return Err(ParseError::MissingHost);
451 }
452 break;
453 }
454 _ => host.push(c),
455 }
456 }
457
458 let port = chars.as_str();
460
461 Ok((host, port.into()))
462}
463
464fn get_database(chars: &mut Chars) -> String {
473 let mut database = String::new();
474 for c in chars {
475 match c {
476 '/' if database.is_empty() => {}
477 '?' => break,
478 _ => database.push(c),
479 }
480 }
481 database
482}
483
484fn get_params(params_string: &str) -> Result<BTreeMap<String, String>, ParseError> {
496 params_string
497 .split('&')
498 .map(|kv| {
499 let mut parts = kv.splitn(2, '=');
500 match (parts.next(), parts.next()) {
501 (Some(key), Some(value)) => Ok((key.to_string(), value.to_string())),
502 _ => Err(ParseError::InvalidParams),
503 }
504 })
505 .collect()
506}
507
508impl DSN {
509 #[must_use]
528 pub fn builder() -> DSNBuilder {
529 DSNBuilder::default()
530 }
531}
532
533#[derive(Clone, Debug, Default)]
579pub struct DSNBuilder {
580 driver: String,
581 username: Option<String>,
582 password: Option<String>,
583 protocol: Option<String>,
584 host: Option<String>,
585 port: Option<u16>,
586 socket: Option<String>,
587 database: Option<String>,
588 params: BTreeMap<String, String>,
589}
590
591impl DSNBuilder {
592 #[must_use]
594 pub fn driver(mut self, driver: impl Into<String>) -> Self {
595 self.driver = driver.into();
596 self
597 }
598
599 #[must_use]
601 pub fn username(mut self, username: impl Into<String>) -> Self {
602 self.username = Some(username.into());
603 self
604 }
605
606 #[must_use]
608 pub fn password(mut self, password: impl Into<String>) -> Self {
609 self.password = Some(password.into());
610 self
611 }
612
613 #[must_use]
615 pub fn host(mut self, host: impl Into<String>) -> Self {
616 self.host = Some(host.into());
617 self.protocol = Some("tcp".to_string());
618 self
619 }
620
621 #[must_use]
623 pub const fn port(mut self, port: u16) -> Self {
624 self.port = Some(port);
625 self
626 }
627
628 #[must_use]
630 pub fn socket(mut self, socket: impl Into<String>) -> Self {
631 self.socket = Some(socket.into());
632 self.protocol = Some("unix".to_string());
633 self
634 }
635
636 #[must_use]
638 pub fn database(mut self, database: impl Into<String>) -> Self {
639 self.database = Some(database.into());
640 self
641 }
642
643 #[must_use]
645 pub fn param(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
646 self.params.insert(key.into(), value.into());
647 self
648 }
649
650 #[must_use]
652 pub fn build(self) -> DSN {
653 let protocol = self.protocol.unwrap_or_else(|| "tcp".to_string());
654
655 let (address, host, socket) = if let Some(socket_path) = self.socket {
656 (socket_path.clone(), None, Some(socket_path))
658 } else {
659 let host_name = self.host.clone().unwrap_or_else(|| "localhost".to_string());
661 let addr = self
662 .port
663 .map_or_else(|| host_name.clone(), |port| format!("{host_name}:{port}"));
664 (addr, Some(host_name), None)
665 };
666
667 DSN {
668 driver: self.driver,
669 username: self.username,
670 password: self.password,
671 protocol,
672 address,
673 host,
674 port: self.port,
675 database: self.database,
676 socket,
677 params: self.params,
678 }
679 }
680}
681
682impl DSNBuilder {
683 #[must_use]
701 pub fn mysql() -> Self {
702 Self {
703 driver: "mysql".to_string(),
704 protocol: Some("tcp".to_string()),
705 port: Some(3306),
706 ..Default::default()
707 }
708 }
709
710 #[must_use]
728 pub fn postgres() -> Self {
729 Self {
730 driver: "postgres".to_string(),
731 protocol: Some("tcp".to_string()),
732 port: Some(5432),
733 ..Default::default()
734 }
735 }
736
737 #[must_use]
754 pub fn redis() -> Self {
755 Self {
756 driver: "redis".to_string(),
757 protocol: Some("tcp".to_string()),
758 port: Some(6379),
759 ..Default::default()
760 }
761 }
762
763 #[must_use]
780 pub fn mariadb() -> Self {
781 Self {
782 driver: "mariadb".to_string(),
783 protocol: Some("tcp".to_string()),
784 port: Some(3306),
785 ..Default::default()
786 }
787 }
788}
789
790#[cfg(test)]
791#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
792mod tests {
793 use super::{DSN, DSNBuilder, ParseError, parse};
794
795 #[test]
796 fn test_parse_password() {
797 let dsn = parse(r#"mysql://user:pas':"'sword44444@host:port/database"#).unwrap();
798 assert_eq!(dsn.password.unwrap(), r#"pas':"'sword44444"#);
799 }
800
801 #[test]
802 fn test_parse_driver() {
803 let dsn = parse(r"mysql://user:pass@host:port/database").unwrap();
804 assert_eq!(dsn.driver, "mysql");
805 }
806
807 #[test]
808 fn test_parse_driver_postgres() {
809 let dsn = parse(r"postgres://user:pass@host:port/database").unwrap();
810 assert_eq!(dsn.driver, "postgres");
811 }
812
813 #[test]
814 fn test_parse_username() {
815 let dsn = parse(r"mysql://user:pass@host:port/database").unwrap();
816 assert_eq!(dsn.username.unwrap(), "user");
817 }
818
819 #[test]
820 fn test_parse_protocol() {
821 let dsn = parse(r"mysql://user:pass@tcp(host:3306)/database").unwrap();
822 assert_eq!(dsn.protocol, "tcp");
823 }
824
825 #[test]
826 fn test_parse_address() {
827 let dsn = parse(r"mysql://user:pass@tcp(host:3306)/database").unwrap();
828 assert_eq!(dsn.address, "host:3306");
829 }
830
831 #[test]
832 fn test_parse_host() {
833 let dsn = parse(r"mysql://user:pass@tcp(host:3306)/database").unwrap();
834 assert_eq!(dsn.host.unwrap(), "host");
835 }
836
837 #[test]
838 fn test_parse_port() {
839 let dsn = parse(r"mysql://user:pass@tcp(host:3306)/database").unwrap();
840 assert_eq!(dsn.port.unwrap(), 3306);
841 }
842
843 #[test]
844 fn test_builder_mysql() {
845 let dsn = DSNBuilder::mysql()
846 .username("root")
847 .password("secret")
848 .host("localhost")
849 .database("mydb")
850 .param("charset", "utf8mb4")
851 .build();
852
853 assert_eq!(dsn.driver, "mysql");
854 assert_eq!(dsn.username.as_deref(), Some("root"));
855 assert_eq!(dsn.password.as_deref(), Some("secret"));
856 assert_eq!(dsn.host.as_deref(), Some("localhost"));
857 assert_eq!(dsn.port, Some(3306));
858 assert_eq!(dsn.database.as_deref(), Some("mydb"));
859 assert_eq!(dsn.params.get("charset"), Some(&"utf8mb4".to_string()));
860 }
861
862 #[test]
863 fn test_builder_postgres() {
864 let dsn = DSNBuilder::postgres()
865 .username("postgres")
866 .password("pass")
867 .host("db.example.com")
868 .database("production")
869 .param("sslmode", "require")
870 .build();
871
872 assert_eq!(dsn.driver, "postgres");
873 assert_eq!(dsn.port, Some(5432));
874 assert_eq!(dsn.params.get("sslmode"), Some(&"require".to_string()));
875 }
876
877 #[test]
878 fn test_builder_redis() {
879 let dsn = DSNBuilder::redis()
880 .host("localhost")
881 .password("secret")
882 .database("0")
883 .build();
884
885 assert_eq!(dsn.driver, "redis");
886 assert_eq!(dsn.port, Some(6379));
887 assert_eq!(dsn.database.as_deref(), Some("0"));
888 }
889
890 #[test]
891 fn test_builder_unix_socket() {
892 let dsn = DSNBuilder::mysql()
893 .username("app")
894 .socket("/var/run/mysqld/mysqld.sock")
895 .database("appdb")
896 .build();
897
898 assert_eq!(dsn.protocol, "unix");
899 assert_eq!(dsn.socket.as_deref(), Some("/var/run/mysqld/mysqld.sock"));
900 assert_eq!(dsn.address, "/var/run/mysqld/mysqld.sock");
901 }
902
903 #[test]
904 fn test_to_string_basic() {
905 let dsn = DSNBuilder::mysql()
906 .username("root")
907 .password("secret")
908 .host("localhost")
909 .database("mydb")
910 .build();
911
912 let dsn_string = dsn.to_string();
913 assert!(dsn_string.contains("mysql://"));
914 assert!(dsn_string.contains("root"));
915 assert!(dsn_string.contains("secret"));
916 assert!(dsn_string.contains("localhost:3306"));
917 assert!(dsn_string.contains("/mydb"));
918 }
919
920 #[test]
921 fn test_to_string_with_params() {
922 let dsn = DSNBuilder::postgres()
923 .username("user")
924 .password("pass")
925 .host("localhost")
926 .database("db")
927 .param("sslmode", "require")
928 .param("connect_timeout", "10")
929 .build();
930
931 let dsn_string = dsn.to_string();
932 assert!(dsn_string.contains('?'));
933 assert!(dsn_string.contains("sslmode=require"));
934 assert!(dsn_string.contains("connect_timeout=10"));
935 }
936
937 #[test]
938 fn test_to_string_special_chars() {
939 let dsn = DSNBuilder::mysql()
940 .username("user@host")
941 .password("p@ss:word!")
942 .host("localhost")
943 .database("mydb")
944 .build();
945
946 let dsn_string = dsn.to_string();
947 assert!(dsn_string.contains("%40")); assert!(!dsn_string.contains("user@host"));
950 }
951
952 #[test]
953 fn test_roundtrip() {
954 let original = "mysql://root:secret@tcp(localhost:3306)/mydb?charset=utf8mb4";
955 let parsed = parse(original).unwrap();
956 let rebuilt = parsed.to_string();
957
958 let reparsed = parse(&rebuilt).unwrap();
960 assert_eq!(parsed.driver, reparsed.driver);
961 assert_eq!(parsed.username, reparsed.username);
962 assert_eq!(parsed.host, reparsed.host);
963 assert_eq!(parsed.port, reparsed.port);
964 assert_eq!(parsed.database, reparsed.database);
965 }
966
967 #[test]
968 fn test_builder_mariadb() {
969 let dsn = DSNBuilder::mariadb()
970 .username("root")
971 .host("localhost")
972 .database("mydb")
973 .build();
974
975 assert_eq!(dsn.driver, "mariadb");
976 assert_eq!(dsn.port, Some(3306));
977 }
978
979 #[test]
980 fn test_error_display() {
981 assert_eq!(format!("{}", ParseError::InvalidDriver), "invalid driver");
983 assert_eq!(format!("{}", ParseError::InvalidParams), "invalid params");
984 assert_eq!(
985 format!("{}", ParseError::InvalidPath),
986 "invalid absolute path"
987 );
988 assert_eq!(
989 format!("{}", ParseError::InvalidPort),
990 "invalid port number"
991 );
992 assert_eq!(
993 format!("{}", ParseError::InvalidProtocol),
994 "invalid protocol"
995 );
996 assert_eq!(format!("{}", ParseError::InvalidSocket), "invalid socket");
997 assert_eq!(format!("{}", ParseError::MissingAddress), "missing address");
998 assert_eq!(format!("{}", ParseError::MissingHost), "missing host");
999 assert_eq!(
1000 format!("{}", ParseError::MissingProtocol),
1001 "missing protocol"
1002 );
1003 assert_eq!(
1004 format!("{}", ParseError::MissingSocket),
1005 "missing unix domain socket"
1006 );
1007 }
1008
1009 #[test]
1010 #[allow(invalid_from_utf8)]
1011 fn test_utf8_error_from() {
1012 let bad_bytes: &[u8] = &[0xFF, 0xFF];
1014 let utf8_err = std::str::from_utf8(bad_bytes).unwrap_err();
1015 let parse_err = ParseError::from(utf8_err);
1016 match parse_err {
1017 ParseError::Utf8Error(_) => {
1018 assert!(format!("{parse_err}").contains("UTF-8 error"));
1019 }
1020 _ => panic!("Expected Utf8Error variant"),
1021 }
1022 }
1023
1024 #[test]
1025 fn test_to_string_no_credentials() {
1026 let dsn = DSNBuilder::mysql().host("localhost").database("db").build();
1028
1029 let dsn_string = dsn.to_string();
1030 assert!(dsn_string.contains("mysql://"));
1031 assert!(!dsn_string.contains('@')); assert!(dsn_string.contains("tcp(localhost:3306)"));
1033 }
1034
1035 #[test]
1036 fn test_to_string_no_database() {
1037 let dsn = DSNBuilder::mysql()
1039 .username("root")
1040 .password("pass")
1041 .host("localhost")
1042 .build();
1043
1044 let dsn_string = dsn.to_string();
1045 assert!(dsn_string.contains("mysql://"));
1046 assert!(dsn_string.ends_with("tcp(localhost:3306)")); }
1048
1049 #[test]
1050 fn test_to_string_username_only() {
1051 let dsn = DSNBuilder::mysql()
1053 .username("root")
1054 .host("localhost")
1055 .database("db")
1056 .build();
1057
1058 let dsn_string = dsn.to_string();
1059 assert!(dsn_string.contains("mysql://root@"));
1060 assert!(!dsn_string.contains(":@")); }
1062
1063 #[test]
1064 fn test_builder_default() {
1065 let dsn = DSNBuilder::default()
1067 .driver("custom")
1068 .host("localhost")
1069 .port(9999)
1070 .build();
1071
1072 assert_eq!(dsn.driver, "custom");
1073 assert_eq!(dsn.port, Some(9999));
1074 }
1075
1076 #[test]
1077 fn test_builder_const_port() {
1078 let dsn = DSNBuilder::mysql().port(3307).host("localhost").build();
1080
1081 assert_eq!(dsn.port, Some(3307));
1082 }
1083
1084 #[test]
1085 fn test_parse_errors() {
1086 assert!(parse("mysql://user@tcp(host:99999)/db").is_err()); assert!(parse("mysql://user@unix(relative/path)/db").is_err()); assert!(parse("mysql://user@file(relative/path)/db").is_err()); assert!(parse("mysql://user@tcp()/db").is_err()); assert!(parse("mysql://user@tcp(:3306)/db").is_err()); assert!(parse("mysql://user@tcp(host:port)/db").is_err()); }
1094
1095 #[test]
1096 fn test_parse_edge_cases() {
1097 let dsn = parse("://user@tcp(host)/db").unwrap();
1099 assert_eq!(dsn.driver, "");
1100
1101 let dsn = parse("mysql://user@udp(host:9999)/db").unwrap();
1103 assert_eq!(dsn.protocol, "udp");
1104 }
1105
1106 #[test]
1107 fn test_parse_missing_protocol() {
1108 assert!(parse("mysql://user@(host)/db").is_err());
1110 }
1111
1112 #[test]
1113 fn test_dsn_builder_method() {
1114 let dsn = DSN::builder().driver("mysql").host("localhost").build();
1116
1117 assert_eq!(dsn.driver, "mysql");
1118 }
1119
1120 #[test]
1121 fn test_dsn_clone() {
1122 let original = parse("mysql://user:pass@tcp(localhost:3306)/mydb?charset=utf8").unwrap();
1123 let cloned = original.clone();
1124
1125 assert_eq!(original.driver, cloned.driver);
1126 assert_eq!(original.username, cloned.username);
1127 assert_eq!(original.password, cloned.password);
1128 assert_eq!(original.protocol, cloned.protocol);
1129 assert_eq!(original.address, cloned.address);
1130 assert_eq!(original.host, cloned.host);
1131 assert_eq!(original.port, cloned.port);
1132 assert_eq!(original.database, cloned.database);
1133 assert_eq!(original.params, cloned.params);
1134 }
1135
1136 #[test]
1137 fn test_dsn_eq() {
1138 let dsn1 = parse("mysql://user:pass@tcp(localhost:3306)/mydb").unwrap();
1139 let dsn2 = parse("mysql://user:pass@tcp(localhost:3306)/mydb").unwrap();
1140 let dsn3 = parse("mysql://user:pass@tcp(localhost:3307)/mydb").unwrap();
1141
1142 assert_eq!(dsn1, dsn2);
1143 assert_ne!(dsn1, dsn3);
1144 }
1145
1146 #[test]
1147 fn test_dsn_hash() {
1148 use std::collections::HashSet;
1149
1150 let dsn1 = parse("mysql://user:pass@tcp(localhost:3306)/mydb").unwrap();
1151 let dsn2 = parse("mysql://user:pass@tcp(localhost:3306)/mydb").unwrap();
1152
1153 let mut set = HashSet::new();
1154 set.insert(dsn1.clone());
1155 set.insert(dsn2);
1156
1157 assert_eq!(set.len(), 1);
1158 assert!(set.contains(&dsn1));
1159 }
1160
1161 #[test]
1162 fn test_dsn_builder_clone() {
1163 let builder1 = DSNBuilder::mysql().username("root").host("localhost");
1164
1165 let builder2 = builder1.clone().database("db1").build();
1166 let builder3 = builder1.database("db2").build();
1167
1168 assert_eq!(builder2.database.as_deref(), Some("db1"));
1169 assert_eq!(builder3.database.as_deref(), Some("db2"));
1170 }
1171}