1use rustc_hash::FxHashSet;
2use std::fmt;
3
4use enum_iterator::Sequence;
5use enum_iterator::all;
6pub use ignore::Ignore;
7use ignore::find_ignores;
8use ignore::has_disable_assume_in_transaction;
9use ignore_index::IgnoreIndex;
10use rowan::TextRange;
11use rowan::TextSize;
12use serde::Deserialize;
13
14use squawk_syntax::SyntaxNode;
15use squawk_syntax::{Parse, SourceFile};
16
17pub use version::Version;
18
19pub mod analyze;
20pub mod ignore;
21mod ignore_index;
22mod version;
23mod visitors;
24
25mod rules;
26
27#[cfg(test)]
28mod test_utils;
29use rules::adding_field_with_default;
30use rules::adding_foreign_key_constraint;
31use rules::adding_not_null_field;
32use rules::adding_primary_key_constraint;
33use rules::adding_required_field;
34use rules::ban_alter_domain_with_add_constraint;
35use rules::ban_char_field;
36use rules::ban_concurrent_index_creation_in_transaction;
37use rules::ban_create_domain_with_constraint;
38use rules::ban_drop_column;
39use rules::ban_drop_database;
40use rules::ban_drop_not_null;
41use rules::ban_drop_table;
42use rules::ban_duplicate_column_assignments;
43use rules::ban_truncate_cascade;
44use rules::ban_uncommitted_transaction;
45use rules::changing_column_type;
46use rules::constraint_missing_not_valid;
47use rules::disallow_unique_constraint;
48use rules::identifier_too_long;
49use rules::prefer_bigint_over_int;
50use rules::prefer_bigint_over_smallint;
51use rules::prefer_identity;
52use rules::prefer_repack;
53use rules::prefer_robust_stmts;
54use rules::prefer_text_field;
55use rules::prefer_timestamptz;
56use rules::renaming_column;
57use rules::renaming_table;
58use rules::require_concurrent_index_creation;
59use rules::require_concurrent_index_deletion;
60use rules::require_concurrent_partition_detach;
61use rules::require_concurrent_reindex;
62use rules::require_enum_value_ordering;
63use rules::require_table_schema;
64use rules::require_timeout_settings;
65use rules::transaction_nesting;
66#[derive(Debug, PartialEq, Clone, Copy, Hash, Eq, Sequence)]
69pub enum Rule {
70 RequireConcurrentIndexCreation,
71 RequireConcurrentIndexDeletion,
72 ConstraintMissingNotValid,
73 AddingFieldWithDefault,
74 AddingForeignKeyConstraint,
75 ChangingColumnType,
76 AddingNotNullableField,
77 AddingSerialPrimaryKeyField,
78 RenamingColumn,
79 RenamingTable,
80 DisallowedUniqueConstraint,
81 BanDropDatabase,
82 PreferBigintOverInt,
83 PreferBigintOverSmallint,
84 PreferIdentity,
85 PreferRepack,
86 PreferRobustStmts,
87 PreferTextField,
88 PreferTimestampTz,
89 BanCharField,
90 BanDropColumn,
91 BanDropTable,
92 BanDropNotNull,
93 TransactionNesting,
94 AddingRequiredField,
95 BanConcurrentIndexCreationInTransaction,
96 UnusedIgnore,
97 BanCreateDomainWithConstraint,
98 BanAlterDomainWithAddConstraint,
99 BanTruncateCascade,
100 RequireTimeoutSettings,
101 BanUncommittedTransaction,
102 RequireEnumValueOrdering,
103 RequireTableSchema,
104 IdentifierTooLong,
105 RequireConcurrentPartitionDetach,
106 RequireConcurrentReindex,
107 RequireLockTimeout,
108 RequireStatementTimeout,
109 BanDuplicateColumnAssignments,
110 }
112
113impl Rule {
114 pub fn is_opt_in(&self) -> bool {
117 matches!(
119 self,
120 Rule::RequireTableSchema | Rule::RequireTimeoutSettings
121 )
122 }
123
124 pub fn expands_to(&self) -> &[Rule] {
126 match self {
127 Rule::RequireTimeoutSettings => {
128 &[Rule::RequireLockTimeout, Rule::RequireStatementTimeout]
129 }
130 _ => &[],
131 }
132 }
133}
134
135impl TryFrom<&str> for Rule {
136 type Error = String;
137
138 fn try_from(s: &str) -> Result<Self, Self::Error> {
139 match s {
140 "require-concurrent-index-creation" => Ok(Rule::RequireConcurrentIndexCreation),
141 "require-concurrent-index-deletion" => Ok(Rule::RequireConcurrentIndexDeletion),
142 "constraint-missing-not-valid" => Ok(Rule::ConstraintMissingNotValid),
143 "adding-field-with-default" => Ok(Rule::AddingFieldWithDefault),
144 "adding-foreign-key-constraint" => Ok(Rule::AddingForeignKeyConstraint),
145 "changing-column-type" => Ok(Rule::ChangingColumnType),
146 "adding-not-nullable-field" => Ok(Rule::AddingNotNullableField),
147 "adding-serial-primary-key-field" => Ok(Rule::AddingSerialPrimaryKeyField),
148 "renaming-column" => Ok(Rule::RenamingColumn),
149 "renaming-table" => Ok(Rule::RenamingTable),
150 "disallowed-unique-constraint" => Ok(Rule::DisallowedUniqueConstraint),
151 "ban-drop-database" => Ok(Rule::BanDropDatabase),
152 "prefer-bigint-over-int" => Ok(Rule::PreferBigintOverInt),
153 "prefer-bigint-over-smallint" => Ok(Rule::PreferBigintOverSmallint),
154 "prefer-identity" => Ok(Rule::PreferIdentity),
155 "prefer-repack" => Ok(Rule::PreferRepack),
156 "prefer-robust-stmts" => Ok(Rule::PreferRobustStmts),
157 "prefer-text-field" => Ok(Rule::PreferTextField),
158 "prefer-timestamptz" => Ok(Rule::PreferTimestampTz),
160 "prefer-timestamp-tz" => Ok(Rule::PreferTimestampTz),
161 "ban-char-field" => Ok(Rule::BanCharField),
162 "ban-drop-column" => Ok(Rule::BanDropColumn),
163 "ban-drop-table" => Ok(Rule::BanDropTable),
164 "ban-drop-not-null" => Ok(Rule::BanDropNotNull),
165 "transaction-nesting" => Ok(Rule::TransactionNesting),
166 "adding-required-field" => Ok(Rule::AddingRequiredField),
167 "ban-concurrent-index-creation-in-transaction" => {
168 Ok(Rule::BanConcurrentIndexCreationInTransaction)
169 }
170 "ban-create-domain-with-constraint" => Ok(Rule::BanCreateDomainWithConstraint),
171 "ban-alter-domain-with-add-constraint" => Ok(Rule::BanAlterDomainWithAddConstraint),
172 "ban-truncate-cascade" => Ok(Rule::BanTruncateCascade),
173 "require-timeout-settings" => Ok(Rule::RequireTimeoutSettings),
174 "ban-uncommitted-transaction" => Ok(Rule::BanUncommittedTransaction),
175 "require-enum-value-ordering" => Ok(Rule::RequireEnumValueOrdering),
176 "require-table-schema" => Ok(Rule::RequireTableSchema),
177 "identifier-too-long" => Ok(Rule::IdentifierTooLong),
178 "require-concurrent-partition-detach" => Ok(Rule::RequireConcurrentPartitionDetach),
179 "require-concurrent-reindex" => Ok(Rule::RequireConcurrentReindex),
180 "require-lock-timeout" => Ok(Rule::RequireLockTimeout),
181 "require-statement-timeout" => Ok(Rule::RequireStatementTimeout),
182 "ban-duplicate-column-assignments" => Ok(Rule::BanDuplicateColumnAssignments),
183 _ => Err(format!("Unknown violation name: {s}")),
185 }
186 }
187}
188
189#[derive(Debug, Clone, PartialEq, Eq)]
190pub struct UnknownRuleName {
191 val: String,
192}
193
194impl std::fmt::Display for UnknownRuleName {
195 fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
196 write!(f, "invalid rule name {}", self.val)
197 }
198}
199
200impl std::error::Error for UnknownRuleName {}
201
202impl std::str::FromStr for Rule {
203 type Err = UnknownRuleName;
204 fn from_str(s: &str) -> Result<Self, Self::Err> {
205 Rule::try_from(s).map_err(|_| UnknownRuleName { val: s.to_string() })
206 }
207}
208
209impl fmt::Display for Rule {
210 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
211 let val = match &self {
212 Rule::RequireConcurrentIndexCreation => "require-concurrent-index-creation",
213 Rule::RequireConcurrentIndexDeletion => "require-concurrent-index-deletion",
214 Rule::ConstraintMissingNotValid => "constraint-missing-not-valid",
215 Rule::AddingFieldWithDefault => "adding-field-with-default",
216 Rule::AddingForeignKeyConstraint => "adding-foreign-key-constraint",
217 Rule::ChangingColumnType => "changing-column-type",
218 Rule::AddingNotNullableField => "adding-not-nullable-field",
219 Rule::AddingSerialPrimaryKeyField => "adding-serial-primary-key-field",
220 Rule::RenamingColumn => "renaming-column",
221 Rule::RenamingTable => "renaming-table",
222 Rule::DisallowedUniqueConstraint => "disallowed-unique-constraint",
223 Rule::BanDropDatabase => "ban-drop-database",
224 Rule::PreferBigintOverInt => "prefer-bigint-over-int",
225 Rule::PreferBigintOverSmallint => "prefer-bigint-over-smallint",
226 Rule::PreferIdentity => "prefer-identity",
227 Rule::PreferRepack => "prefer-repack",
228 Rule::PreferRobustStmts => "prefer-robust-stmts",
229 Rule::PreferTextField => "prefer-text-field",
230 Rule::PreferTimestampTz => "prefer-timestamp-tz",
231 Rule::BanCharField => "ban-char-field",
232 Rule::BanDropColumn => "ban-drop-column",
233 Rule::BanDropTable => "ban-drop-table",
234 Rule::BanDropNotNull => "ban-drop-not-null",
235 Rule::TransactionNesting => "transaction-nesting",
236 Rule::AddingRequiredField => "adding-required-field",
237 Rule::BanConcurrentIndexCreationInTransaction => {
238 "ban-concurrent-index-creation-in-transaction"
239 }
240 Rule::BanCreateDomainWithConstraint => "ban-create-domain-with-constraint",
241 Rule::UnusedIgnore => "unused-ignore",
242 Rule::BanAlterDomainWithAddConstraint => "ban-alter-domain-with-add-constraint",
243 Rule::BanTruncateCascade => "ban-truncate-cascade",
244 Rule::RequireTimeoutSettings => "require-timeout-settings",
245 Rule::BanUncommittedTransaction => "ban-uncommitted-transaction",
246 Rule::RequireEnumValueOrdering => "require-enum-value-ordering",
247 Rule::RequireTableSchema => "require-table-schema",
248 Rule::IdentifierTooLong => "identifier-too-long",
249 Rule::RequireConcurrentPartitionDetach => "require-concurrent-partition-detach",
250 Rule::RequireConcurrentReindex => "require-concurrent-reindex",
251 Rule::RequireLockTimeout => "require-lock-timeout",
252 Rule::RequireStatementTimeout => "require-statement-timeout",
253 Rule::BanDuplicateColumnAssignments => "ban-duplicate-column-assignments",
254 };
256 write!(f, "{val}")
257 }
258}
259
260impl<'de> Deserialize<'de> for Rule {
261 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
262 where
263 D: serde::Deserializer<'de>,
264 {
265 let s = String::deserialize(deserializer)?;
266 s.parse().map_err(serde::de::Error::custom)
267 }
268}
269
270#[derive(Debug, Clone, PartialEq, Eq)]
271pub struct Fix {
272 pub title: String,
273 pub edits: Vec<Edit>,
274}
275
276impl Fix {
277 fn new<T: Into<String>>(title: T, edits: Vec<Edit>) -> Fix {
278 Fix {
279 title: title.into(),
280 edits,
281 }
282 }
283}
284
285#[derive(Debug, Clone, PartialEq, Eq)]
286pub struct Edit {
287 pub text_range: TextRange,
288 pub text: Option<String>,
290}
291impl Edit {
292 pub fn insert<T: Into<String>>(text: T, at: TextSize) -> Self {
293 Self {
294 text_range: TextRange::new(at, at),
295 text: Some(text.into()),
296 }
297 }
298 pub fn replace<T: Into<String>>(text_range: TextRange, text: T) -> Self {
299 Self {
300 text_range,
301 text: Some(text.into()),
302 }
303 }
304 pub fn delete(text_range: TextRange) -> Self {
305 Self {
306 text_range,
307 text: None,
308 }
309 }
310}
311
312#[derive(Debug, Clone, PartialEq, Eq)]
313pub struct Violation {
314 pub code: Rule,
316 pub message: String,
317 pub text_range: TextRange,
318 pub help: Option<String>,
319 pub fix: Option<Fix>,
320}
321
322impl Violation {
323 #[must_use]
324 pub fn for_node(code: Rule, message: String, node: &SyntaxNode) -> Self {
325 let range = node.text_range();
326
327 let start = node
328 .children_with_tokens()
329 .find(|x| !x.kind().is_trivia())
330 .map(|x| x.text_range().start())
331 .unwrap_or_else(|| range.start());
333
334 Self {
335 code,
336 text_range: TextRange::new(start, range.end()),
337 message,
338 help: None,
339 fix: None,
340 }
341 }
342
343 #[must_use]
344 pub fn for_range(code: Rule, message: String, text_range: TextRange) -> Self {
345 Self {
346 code,
347 text_range,
348 message,
349 help: None,
350 fix: None,
351 }
352 }
353
354 fn fix<F: Into<Option<Fix>>>(mut self, fix: F) -> Violation {
355 self.fix = fix.into();
356 self
357 }
358 fn help(mut self, help: impl Into<String>) -> Violation {
359 self.help = Some(help.into());
360 self
361 }
362}
363
364#[derive(Clone, Default)]
365pub struct LinterSettings {
366 pub pg_version: Version,
367 pub assume_in_transaction: bool,
368}
369
370pub struct Linter {
371 errors: Vec<Violation>,
372 ignores: Vec<Ignore>,
373 pub rules: FxHashSet<Rule>,
374 pub settings: LinterSettings,
375}
376
377impl Linter {
378 fn report(&mut self, error: Violation) {
379 self.errors.push(error);
380 }
381
382 fn ignore(&mut self, ignore: Ignore) {
383 self.ignores.push(ignore);
384 }
385
386 #[must_use]
387 pub fn lint(&mut self, file: &Parse<SourceFile>, text: &str) -> Vec<Violation> {
388 if has_disable_assume_in_transaction(&file.syntax_node()) {
389 self.settings.assume_in_transaction = false;
390 }
391
392 if self.rules.contains(&Rule::AddingFieldWithDefault) {
393 adding_field_with_default(self, file);
394 }
395 if self.rules.contains(&Rule::AddingForeignKeyConstraint) {
396 adding_foreign_key_constraint(self, file);
397 }
398 if self.rules.contains(&Rule::AddingNotNullableField) {
399 adding_not_null_field(self, file);
400 }
401 if self.rules.contains(&Rule::AddingSerialPrimaryKeyField) {
402 adding_primary_key_constraint(self, file);
403 }
404 if self.rules.contains(&Rule::AddingRequiredField) {
405 adding_required_field(self, file);
406 }
407 if self.rules.contains(&Rule::BanDropDatabase) {
408 ban_drop_database(self, file);
409 }
410 if self.rules.contains(&Rule::BanCharField) {
411 ban_char_field(self, file);
412 }
413 if self
414 .rules
415 .contains(&Rule::BanConcurrentIndexCreationInTransaction)
416 {
417 ban_concurrent_index_creation_in_transaction(self, file);
418 }
419 if self.rules.contains(&Rule::BanDropColumn) {
420 ban_drop_column(self, file);
421 }
422 if self.rules.contains(&Rule::BanDropNotNull) {
423 ban_drop_not_null(self, file);
424 }
425 if self.rules.contains(&Rule::BanDropTable) {
426 ban_drop_table(self, file);
427 }
428 if self.rules.contains(&Rule::ChangingColumnType) {
429 changing_column_type(self, file);
430 }
431 if self.rules.contains(&Rule::ConstraintMissingNotValid) {
432 constraint_missing_not_valid(self, file);
433 }
434 if self.rules.contains(&Rule::DisallowedUniqueConstraint) {
435 disallow_unique_constraint(self, file);
436 }
437 if self.rules.contains(&Rule::PreferBigintOverInt) {
438 prefer_bigint_over_int(self, file);
439 }
440 if self.rules.contains(&Rule::PreferBigintOverSmallint) {
441 prefer_bigint_over_smallint(self, file);
442 }
443 if self.rules.contains(&Rule::PreferIdentity) {
444 prefer_identity(self, file);
445 }
446 if self.rules.contains(&Rule::PreferRepack) {
447 prefer_repack(self, file);
448 }
449 if self.rules.contains(&Rule::PreferRobustStmts) {
450 prefer_robust_stmts(self, file);
451 }
452 if self.rules.contains(&Rule::PreferTextField) {
453 prefer_text_field(self, file);
454 }
455 if self.rules.contains(&Rule::PreferTimestampTz) {
456 prefer_timestamptz(self, file);
457 }
458 if self.rules.contains(&Rule::RenamingColumn) {
459 renaming_column(self, file);
460 }
461 if self.rules.contains(&Rule::RenamingTable) {
462 renaming_table(self, file);
463 }
464 if self.rules.contains(&Rule::RequireConcurrentIndexCreation) {
465 require_concurrent_index_creation(self, file);
466 }
467 if self.rules.contains(&Rule::RequireConcurrentIndexDeletion) {
468 require_concurrent_index_deletion(self, file);
469 }
470 if self.rules.contains(&Rule::BanCreateDomainWithConstraint) {
471 ban_create_domain_with_constraint(self, file);
472 }
473 if self.rules.contains(&Rule::BanAlterDomainWithAddConstraint) {
474 ban_alter_domain_with_add_constraint(self, file);
475 }
476 if self.rules.contains(&Rule::TransactionNesting) {
477 transaction_nesting(self, file);
478 }
479 if self.rules.contains(&Rule::BanTruncateCascade) {
480 ban_truncate_cascade(self, file);
481 }
482 if self.rules.contains(&Rule::RequireLockTimeout)
483 || self.rules.contains(&Rule::RequireStatementTimeout)
484 {
485 require_timeout_settings(self, file);
486 }
487 if self.rules.contains(&Rule::BanUncommittedTransaction) {
488 ban_uncommitted_transaction(self, file);
489 }
490 if self.rules.contains(&Rule::RequireEnumValueOrdering) {
491 require_enum_value_ordering(self, file);
492 }
493 if self.rules.contains(&Rule::RequireTableSchema) {
494 require_table_schema(self, file);
495 }
496 if self.rules.contains(&Rule::IdentifierTooLong) {
497 identifier_too_long(self, file);
498 }
499 if self.rules.contains(&Rule::RequireConcurrentPartitionDetach) {
500 require_concurrent_partition_detach(self, file);
501 }
502 if self.rules.contains(&Rule::RequireConcurrentReindex) {
503 require_concurrent_reindex(self, file);
504 }
505 if self.rules.contains(&Rule::BanDuplicateColumnAssignments) {
506 ban_duplicate_column_assignments(self, file);
507 }
508 find_ignores(self, &file.syntax_node());
512
513 self.errors(text)
514 }
515
516 fn errors(&mut self, text: &str) -> Vec<Violation> {
517 let ignore_index = IgnoreIndex::new(text, &self.ignores);
518 let mut errors: Vec<Violation> = self
519 .errors
520 .iter()
521 .filter(|err| !ignore_index.contains(err.text_range, err.code))
524 .cloned()
525 .collect::<Vec<_>>();
526 errors.sort_by_key(|x| x.text_range.start());
528 errors
529 }
530
531 fn default_rules() -> FxHashSet<Rule> {
532 all::<Rule>()
533 .filter(|r| !r.is_opt_in())
534 .collect::<FxHashSet<_>>()
535 }
536
537 pub fn with_default_rules() -> Self {
538 let rules = Linter::default_rules();
539 Linter::from(rules)
540 }
541
542 pub fn with_rules(include: &[Rule], exclude: &[Rule]) -> Self {
543 let mut default_rules = Linter::default_rules();
544
545 for rule in include {
546 default_rules.insert(*rule);
547 default_rules.extend(rule.expands_to());
548 }
549
550 for rule in exclude {
551 default_rules.remove(rule);
552 for expanded in rule.expands_to() {
553 default_rules.remove(expanded);
554 }
555 }
556
557 default_rules.retain(|rule| rule.expands_to().is_empty());
560
561 Linter::from(default_rules)
562 }
563
564 pub fn from(rules: impl IntoIterator<Item = Rule>) -> Self {
565 let mut rules: FxHashSet<Rule> = rules.into_iter().collect();
566 for rule in rules.clone() {
567 rules.extend(rule.expands_to());
568 }
569 Self {
570 errors: vec![],
571 ignores: vec![],
572 rules,
573 settings: Default::default(),
574 }
575 }
576}
577
578#[cfg(test)]
579mod tests {
580 use insta::assert_debug_snapshot;
581
582 use super::*;
583
584 #[test]
585 fn prefer_timestamp_aliases() {
586 let rule1: Rule = "prefer-timestamp-tz".parse().unwrap();
587 let rule2: Rule = "prefer-timestamptz".parse().unwrap();
588 assert_eq!(rule1, rule2);
589 assert_debug_snapshot!(rule1, @"PreferTimestampTz");
590 }
591
592 #[test]
593 fn invalid_rule_name() {
594 let result: Result<Rule, _> = "invalid-rule-name".parse();
595 assert!(result.is_err());
596 }
597
598 #[test]
599 fn with_rules_opt_in_disabled_by_default() {
600 let linter = Linter::with_rules(&[], &[]);
601 assert!(!linter.rules.contains(&Rule::RequireTableSchema));
602 }
603
604 #[test]
605 fn with_rules_opt_in_enabled_via_include() {
606 let linter = Linter::with_rules(&[Rule::RequireTableSchema], &[]);
607 assert!(linter.rules.contains(&Rule::RequireTableSchema));
608 }
609
610 #[test]
611 fn with_rules_exclude_takes_precedence_over_include() {
612 let linter = Linter::with_rules(&[Rule::RequireTableSchema], &[Rule::RequireTableSchema]);
613 assert!(!linter.rules.contains(&Rule::RequireTableSchema));
614 }
615
616 #[test]
617 fn with_rules_exclude_removes_default_rule() {
618 let linter = Linter::with_rules(&[], &[Rule::BanDropTable]);
619 assert!(!linter.rules.contains(&Rule::BanDropTable));
620 }
621
622 #[test]
623 fn require_timeout_settings_expands_to_granular_rules() {
624 let linter = Linter::from([Rule::RequireTimeoutSettings]);
625 assert!(linter.rules.contains(&Rule::RequireLockTimeout));
626 assert!(linter.rules.contains(&Rule::RequireStatementTimeout));
627 }
628
629 #[test]
630 fn with_rules_exclude_timeout_settings_removes_granular_rules() {
631 let linter = Linter::with_rules(&[], &[Rule::RequireTimeoutSettings]);
632 assert!(!linter.rules.contains(&Rule::RequireLockTimeout));
633 assert!(!linter.rules.contains(&Rule::RequireStatementTimeout));
634 }
635
636 #[test]
637 fn with_rules_exclude_granular_timeout_rule_keeps_other() {
638 let linter = Linter::with_rules(&[], &[Rule::RequireStatementTimeout]);
639 assert!(linter.rules.contains(&Rule::RequireLockTimeout));
640 assert!(!linter.rules.contains(&Rule::RequireStatementTimeout));
641 }
642
643 #[test]
644 fn with_rules_exclude_granular_rule_wins_over_included_alias() {
645 let linter = Linter::with_rules(
646 &[Rule::RequireTimeoutSettings],
647 &[Rule::RequireStatementTimeout],
648 );
649 assert!(linter.rules.contains(&Rule::RequireLockTimeout));
650 assert!(!linter.rules.contains(&Rule::RequireStatementTimeout));
651 }
652}