diff --git a/changelog.d/unreleased/2101.fixed.md b/changelog.d/unreleased/2101.fixed.md new file mode 100644 index 0000000000..433548bc8b --- /dev/null +++ b/changelog.d/unreleased/2101.fixed.md @@ -0,0 +1,18 @@ +--- +category: fixed +issues: + - 2101 +affected: + - src/CodeIndex/Database/DbContext.cs + - src/CodeIndex/Database/DbSymbolReader.cs + - src/CodeIndex/Indexer/References/Support/SqlNameResolver.cs + - tests/CodeIndex.Tests/ReferenceExtractorTests.cs +--- + +## English + +- **SQL quoted qualified names now respect dialect matching rules (#2101)** — schema-qualified SQL references keep MySQL backtick and T-SQL bracket matching case-insensitive while preserving PostgreSQL double-quoted identifier case sensitivity. + +## 日本語 + +- **SQL の引用付き qualified name が dialect ごとの照合規則を尊重するようになりました (#2101)** — schema 修飾された SQL 参照では MySQL のバッククォートと T-SQL のブラケットは大文字小文字を区別しない照合を維持しつつ、PostgreSQL の二重引用符付き識別子は大文字小文字を区別します。 diff --git a/src/CodeIndex/Database/DbContext.cs b/src/CodeIndex/Database/DbContext.cs index 1b8d3b7e61..e0f0c6160f 100644 --- a/src/CodeIndex/Database/DbContext.cs +++ b/src/CodeIndex/Database/DbContext.cs @@ -647,6 +647,10 @@ internal static void RegisterConnectionFunctions(SqliteConnection connection) && segmentCount > 0 ? segmentCount : null)); + connection.CreateFunction( + "sql_reference_matches_target_at", + (string? symbolName, string? context, string? containerName, long? columnNumber, string? targetName) => + SqlNameResolver.ReferenceMatchesTargetAtColumn(symbolName, context, containerName, ToNullableInt(columnNumber), targetName) ? 1 : 0); connection.CreateFunction( "sql_allow_leaf_fallback_at", (string? symbolName, string? context, string? containerName, long? columnNumber) => diff --git a/src/CodeIndex/Database/DbSymbolReader.cs b/src/CodeIndex/Database/DbSymbolReader.cs index 035ab301f2..4cae8157ad 100644 --- a/src/CodeIndex/Database/DbSymbolReader.cs +++ b/src/CodeIndex/Database/DbSymbolReader.cs @@ -2570,7 +2570,7 @@ FROM symbol_references sr WHERE sr.symbol_name = s.name OR (f.lang = 'sql' AND rf.lang = 'sql' AND ( (sql_resolve_reference_segment_count_at(sr.symbol_name, " + ReferenceContextSql("sr") + @", sr.container_name, sr.column_number) = sql_segment_count(s.name) - AND sql_resolve_reference_name_at(sr.symbol_name, " + ReferenceContextSql("sr") + @", sr.container_name, sr.column_number) = sql_normalize_name(s.name) COLLATE NOCASE) + AND sql_reference_matches_target_at(sr.symbol_name, " + ReferenceContextSql("sr") + @", sr.container_name, sr.column_number, s.name) = 1) OR (sql_segment_count(sr.symbol_name) = 1 AND sql_allow_leaf_fallback_at(sr.symbol_name, " + ReferenceContextSql("sr") + @", sr.container_name, sr.column_number) = 1 AND sr.symbol_name = sql_leaf_name(s.name) COLLATE NOCASE @@ -2580,7 +2580,7 @@ FROM symbols s_exact JOIN files f_exact ON f_exact.id = s_exact.file_id WHERE f_exact.lang = 'sql' AND sql_segment_count(s_exact.name) = sql_resolve_reference_segment_count_at(sr.symbol_name, " + ReferenceContextSql("sr") + @", sr.container_name, sr.column_number) - AND sql_normalize_name(s_exact.name) = sql_resolve_reference_name_at(sr.symbol_name, " + ReferenceContextSql("sr") + @", sr.container_name, sr.column_number) COLLATE NOCASE + AND sql_reference_matches_target_at(sr.symbol_name, " + ReferenceContextSql("sr") + @", sr.container_name, sr.column_number, s_exact.name) = 1 )) )) )"; @@ -2726,7 +2726,7 @@ FROM symbol_references sr WHERE sr.symbol_name = s.name OR (f.lang = 'sql' AND rf.lang = 'sql' AND ( (sql_resolve_reference_segment_count_at(sr.symbol_name, " + contextSql + @", sr.container_name, sr.column_number) = sql_segment_count(s.name) - AND sql_resolve_reference_name_at(sr.symbol_name, " + contextSql + @", sr.container_name, sr.column_number) = sql_normalize_name(s.name) COLLATE NOCASE) + AND sql_reference_matches_target_at(sr.symbol_name, " + contextSql + @", sr.container_name, sr.column_number, s.name) = 1) OR (sql_segment_count(sr.symbol_name) = 1 AND sql_allow_leaf_fallback_at(sr.symbol_name, " + contextSql + @", sr.container_name, sr.column_number) = 1 AND sr.symbol_name = sql_leaf_name(s.name) COLLATE NOCASE @@ -2736,7 +2736,7 @@ FROM symbols s_exact JOIN files f_exact ON f_exact.id = s_exact.file_id WHERE f_exact.lang = 'sql' AND sql_segment_count(s_exact.name) = sql_resolve_reference_segment_count_at(sr.symbol_name, " + contextSql + @", sr.container_name, sr.column_number) - AND sql_normalize_name(s_exact.name) = sql_resolve_reference_name_at(sr.symbol_name, " + contextSql + @", sr.container_name, sr.column_number) COLLATE NOCASE + AND sql_reference_matches_target_at(sr.symbol_name, " + contextSql + @", sr.container_name, sr.column_number, s_exact.name) = 1 )) )) )"; diff --git a/src/CodeIndex/Indexer/References/Support/SqlNameResolver.cs b/src/CodeIndex/Indexer/References/Support/SqlNameResolver.cs index c979e4afc4..1444afeeba 100644 --- a/src/CodeIndex/Indexer/References/Support/SqlNameResolver.cs +++ b/src/CodeIndex/Indexer/References/Support/SqlNameResolver.cs @@ -5,10 +5,17 @@ namespace CodeIndex.Indexer; internal static class SqlNameResolver { - private readonly record struct SqlNameParts(string NormalizedName, string LeafName, int SegmentCount); + private readonly record struct SqlNameParts( + string NormalizedName, + string LeafName, + int SegmentCount, + IReadOnlyList Segments, + IReadOnlyList CaseSensitiveSegments); private readonly record struct QualifiedNameMatch( string NormalizedName, int SegmentCount, + IReadOnlyList Segments, + IReadOnlyList CaseSensitiveSegments, int StartIndex, int EndIndexExclusive, int LeafStartIndex, @@ -34,7 +41,7 @@ public static bool ContextContainsQualifiedNameAtColumn(string? context, string? return TryGetQualifiedNameAtColumn(context, columnNumber, out var match) && match.SegmentCount == queryParts.SegmentCount - && string.Equals(match.NormalizedName, queryParts.NormalizedName, StringComparison.OrdinalIgnoreCase); + && QualifiedNamesEqual(match.Segments, match.CaseSensitiveSegments, queryParts.Segments, queryParts.CaseSensitiveSegments); } public static bool ContextContainsQualifiedNameLikeAtColumn(string? context, string? query, int? columnNumber) @@ -45,7 +52,7 @@ public static bool ContextContainsQualifiedNameLikeAtColumn(string? context, str return TryGetQualifiedNameAtColumn(context, columnNumber, out var match) && match.SegmentCount == queryParts.SegmentCount - && string.Equals(match.NormalizedName, queryParts.NormalizedName, StringComparison.OrdinalIgnoreCase); + && QualifiedNamesEqual(match.Segments, match.CaseSensitiveSegments, queryParts.Segments, queryParts.CaseSensitiveSegments); } public static bool ContextContainsQualifiedName(string? context, string? query) @@ -57,7 +64,7 @@ public static bool ContextContainsQualifiedName(string? context, string? query) foreach (var candidate in EnumerateQualifiedNameMatches(context)) { if (candidate.SegmentCount == queryParts.SegmentCount - && string.Equals(candidate.NormalizedName, queryParts.NormalizedName, StringComparison.OrdinalIgnoreCase)) + && QualifiedNamesEqual(candidate.Segments, candidate.CaseSensitiveSegments, queryParts.Segments, queryParts.CaseSensitiveSegments)) { return true; } @@ -114,7 +121,7 @@ public static string ResolveReferenceNameAtColumn(string? symbolName, string? co if (leafName.Length > 0 && columnNumber.HasValue && columnNumber.Value > 0) { if (TryGetQualifiedNameAtColumn(context, columnNumber, out var match) - && string.Equals(GetLeafName(match.NormalizedName), leafName, StringComparison.OrdinalIgnoreCase)) + && LeafNamesEqual(match.Segments, match.CaseSensitiveSegments, ParseParts(symbolName))) { return match.NormalizedName; } @@ -160,6 +167,28 @@ public static int ResolveReferenceSegmentCountAtColumn(string? symbolName, strin return GetSegmentCount(ResolveReferenceName(symbolName, context, containerName)); } + public static bool ReferenceMatchesTargetAtColumn( + string? symbolName, + string? context, + string? containerName, + int? columnNumber, + string? targetName) + { + var targetParts = ParseParts(targetName); + if (targetParts.NormalizedName.Length == 0) + return false; + + var resolved = ResolveReferenceNameAtColumn(symbolName, context, containerName, columnNumber); + if (resolved.Length == 0) + return false; + + if (TryGetQualifiedNameAtColumn(context, columnNumber, out var match)) + return QualifiedNamesEqual(match.Segments, match.CaseSensitiveSegments, targetParts.Segments, targetParts.CaseSensitiveSegments); + + var resolvedParts = ParseParts(resolved); + return QualifiedNamesEqual(resolvedParts.Segments, resolvedParts.CaseSensitiveSegments, targetParts.Segments, targetParts.CaseSensitiveSegments); + } + public static bool AllowLeafFallbackAtColumn(string? symbolName, string? context, string? containerName, int? columnNumber) { var normalizedSymbolName = NormalizeQualifiedName(symbolName); @@ -169,7 +198,7 @@ public static bool AllowLeafFallbackAtColumn(string? symbolName, string? context var leafName = GetLeafName(symbolName); if (leafName.Length > 0 && TryGetQualifiedNameAtColumn(context, columnNumber, out var match) - && string.Equals(GetLeafName(match.NormalizedName), leafName, StringComparison.OrdinalIgnoreCase)) + && LeafNamesEqual(match.Segments, match.CaseSensitiveSegments, ParseParts(symbolName))) { return false; } @@ -194,10 +223,8 @@ public static bool ContextContainsQualifiedNameFoldedAtColumn(string? context, s if (!TryGetQualifiedNameAtColumn(context, columnNumber, out var match)) return false; - var foldedCandidate = NameFold.Fold(match.NormalizedName) ?? match.NormalizedName; - var foldedQuery = NameFold.Fold(queryParts.NormalizedName) ?? queryParts.NormalizedName; return match.SegmentCount == queryParts.SegmentCount - && string.Equals(foldedCandidate, foldedQuery, StringComparison.Ordinal); + && QualifiedNamesEqualFolded(match.Segments, match.CaseSensitiveSegments, queryParts.Segments, queryParts.CaseSensitiveSegments); } public static bool ContextContainsQualifiedNameLikeFoldedAtColumn(string? context, string? query, int? columnNumber) @@ -208,10 +235,8 @@ public static bool ContextContainsQualifiedNameLikeFoldedAtColumn(string? contex if (!TryGetQualifiedNameAtColumn(context, columnNumber, out var match)) return false; - var foldedCandidate = NameFold.Fold(match.NormalizedName) ?? match.NormalizedName; - var foldedQuery = NameFold.Fold(queryParts.NormalizedName) ?? queryParts.NormalizedName; return match.SegmentCount == queryParts.SegmentCount - && string.Equals(foldedCandidate, foldedQuery, StringComparison.Ordinal); + && QualifiedNamesEqualFolded(match.Segments, match.CaseSensitiveSegments, queryParts.Segments, queryParts.CaseSensitiveSegments); } public static bool ContextContainsQualifiedNameFolded(string? context, string? query) @@ -220,12 +245,10 @@ public static bool ContextContainsQualifiedNameFolded(string? context, string? q if (queryParts.NormalizedName.Length == 0 || queryParts.SegmentCount <= 1 || string.IsNullOrWhiteSpace(context)) return false; - var foldedQuery = NameFold.Fold(queryParts.NormalizedName) ?? queryParts.NormalizedName; foreach (var candidate in EnumerateQualifiedNameMatches(context)) { - var foldedCandidate = NameFold.Fold(candidate.NormalizedName) ?? candidate.NormalizedName; if (candidate.SegmentCount == queryParts.SegmentCount - && string.Equals(foldedCandidate, foldedQuery, StringComparison.Ordinal)) + && QualifiedNamesEqualFolded(candidate.Segments, candidate.CaseSensitiveSegments, queryParts.Segments, queryParts.CaseSensitiveSegments)) { return true; } @@ -266,12 +289,14 @@ private static bool TryGetQualifiedNameAtColumn(string? context, int? columnNumb private static SqlNameParts ParseParts(string? qualifiedName) { if (string.IsNullOrWhiteSpace(qualifiedName)) - return new SqlNameParts(string.Empty, string.Empty, 0); + return new SqlNameParts(string.Empty, string.Empty, 0, [], []); var trimmed = qualifiedName.Trim(); var segments = new List(); + var caseSensitiveSegments = new List(); var current = new StringBuilder(); char quote = '\0'; + var currentHasCaseSensitiveQuote = false; for (var i = 0; i < trimmed.Length; i++) { @@ -323,25 +348,27 @@ private static SqlNameParts ParseParts(string? qualifiedName) if (ch is '[' or '"' or '`') { quote = ch; + currentHasCaseSensitiveQuote |= ch == '"'; continue; } if (ch == '.') { - AppendNormalizedSegment(segments, current); + AppendNormalizedSegment(segments, caseSensitiveSegments, current, currentHasCaseSensitiveQuote); + currentHasCaseSensitiveQuote = false; continue; } current.Append(ch); } - AppendNormalizedSegment(segments, current); + AppendNormalizedSegment(segments, caseSensitiveSegments, current, currentHasCaseSensitiveQuote); var segmentCount = segments.Count; if (segmentCount == 0) - return new SqlNameParts(string.Empty, string.Empty, 0); + return new SqlNameParts(string.Empty, string.Empty, 0, [], []); var normalized = string.Join(".", segments); - return new SqlNameParts(normalized, segments[^1], segmentCount); + return new SqlNameParts(normalized, segments[^1], segmentCount, segments, caseSensitiveSegments); } private static IEnumerable EnumerateQualifiedNames(string text) @@ -385,11 +412,14 @@ private static string QualifyLeafNameFromContainerCore(string normalizedSymbolNa return containerParts.NormalizedName[..(lastDot + 1)] + normalizedSymbolName; } - private static void AppendNormalizedSegment(List segments, StringBuilder current) + private static void AppendNormalizedSegment(List segments, List caseSensitiveSegments, StringBuilder current, bool hasCaseSensitiveQuote) { var value = current.ToString().Trim(); if (value.Length > 0) + { segments.Add(value); + caseSensitiveSegments.Add(hasCaseSensitiveQuote); + } current.Clear(); } @@ -417,16 +447,18 @@ private static bool TryReadQualifiedName(string text, int startIndex, out Qualif match = default; var segments = new List(); + var caseSensitiveSegments = new List(); var index = startIndex; var leafStartIndex = startIndex; var leafEndIndexExclusive = startIndex; while (true) { var segmentStartIndex = index; - if (!TryReadQualifiedNameSegment(text, ref index, out var segment)) + if (!TryReadQualifiedNameSegment(text, ref index, out var segment, out var segmentHasCaseSensitiveQuote)) return false; segments.Add(segment); + caseSensitiveSegments.Add(segmentHasCaseSensitiveQuote); leafStartIndex = segmentStartIndex; leafEndIndexExclusive = index; var scan = index; @@ -454,13 +486,14 @@ private static bool TryReadQualifiedName(string text, int startIndex, out Qualif if (normalizedName.Length == 0) return false; - match = new QualifiedNameMatch(normalizedName, segments.Count, startIndex, index, leafStartIndex, leafEndIndexExclusive); + match = new QualifiedNameMatch(normalizedName, segments.Count, segments, caseSensitiveSegments, startIndex, index, leafStartIndex, leafEndIndexExclusive); return true; } - private static bool TryReadQualifiedNameSegment(string text, ref int index, out string segment) + private static bool TryReadQualifiedNameSegment(string text, ref int index, out string segment, out bool hasCaseSensitiveQuote) { segment = string.Empty; + hasCaseSensitiveQuote = false; if (index >= text.Length) return false; @@ -468,6 +501,7 @@ private static bool TryReadQualifiedNameSegment(string text, ref int index, out var quote = text[index]; if (quote is '[' or '"' or '`') { + hasCaseSensitiveQuote = quote == '"'; index++; while (index < text.Length) { @@ -528,15 +562,15 @@ private static bool TryReadQualifiedNamePrefixAtColumn( out QualifiedNameMatch match) { match = default; - var segments = new List<(string Name, int StartIndex, int EndIndexExclusive)>(); + var segments = new List<(string Name, bool HasCaseSensitiveQuote, int StartIndex, int EndIndexExclusive)>(); var index = startIndex; while (true) { var segmentStartIndex = index; - if (!TryReadQualifiedNameSegment(text, ref index, out var segment)) + if (!TryReadQualifiedNameSegment(text, ref index, out var segment, out var segmentHasCaseSensitiveQuote)) return false; - segments.Add((segment, segmentStartIndex, index)); + segments.Add((segment, segmentHasCaseSensitiveQuote, segmentStartIndex, index)); var scan = index; while (scan < text.Length && char.IsWhiteSpace(text[scan])) scan++; @@ -560,10 +594,13 @@ private static bool TryReadQualifiedNamePrefixAtColumn( if (zeroBasedColumn < segment.StartIndex || zeroBasedColumn >= segment.EndIndexExclusive) continue; - var normalizedName = string.Join(".", segments.Take(i + 1).Select(part => part.Name)); + var matchedSegments = segments.Take(i + 1).ToList(); + var normalizedName = string.Join(".", matchedSegments.Select(part => part.Name)); match = new QualifiedNameMatch( normalizedName, i + 1, + matchedSegments.Select(part => part.Name).ToList(), + matchedSegments.Select(part => part.HasCaseSensitiveQuote).ToList(), startIndex, segment.EndIndexExclusive, segment.StartIndex, @@ -581,4 +618,65 @@ private static bool IsSqlIdentifierStartChar(char ch) private static bool IsSqlIdentifierChar(char ch) => ch is '_' or '$' or '#' || char.IsLetterOrDigit(ch); + + private static bool LeafNamesEqual( + IReadOnlyList leftSegments, + IReadOnlyList leftCaseSensitiveSegments, + SqlNameParts rightParts) + => rightParts.SegmentCount > 0 + && leftSegments.Count > 0 + && SegmentsEqual( + leftSegments[^1], + leftCaseSensitiveSegments.Count > 0 && leftCaseSensitiveSegments[^1], + rightParts.LeafName, + rightParts.CaseSensitiveSegments.Count > 0 && rightParts.CaseSensitiveSegments[^1]); + + private static bool QualifiedNamesEqual( + IReadOnlyList leftSegments, + IReadOnlyList leftCaseSensitiveSegments, + IReadOnlyList rightSegments, + IReadOnlyList rightCaseSensitiveSegments) + { + if (leftSegments.Count != rightSegments.Count) + return false; + + for (var i = 0; i < leftSegments.Count; i++) + { + if (!SegmentsEqual( + leftSegments[i], + i < leftCaseSensitiveSegments.Count && leftCaseSensitiveSegments[i], + rightSegments[i], + i < rightCaseSensitiveSegments.Count && rightCaseSensitiveSegments[i])) + { + return false; + } + } + + return true; + } + + private static bool QualifiedNamesEqualFolded( + IReadOnlyList leftSegments, + IReadOnlyList leftCaseSensitiveSegments, + IReadOnlyList rightSegments, + IReadOnlyList rightCaseSensitiveSegments) + { + if (leftSegments.Count != rightSegments.Count) + return false; + + for (var i = 0; i < leftSegments.Count; i++) + { + var preserveCase = (i < leftCaseSensitiveSegments.Count && leftCaseSensitiveSegments[i]) + || (i < rightCaseSensitiveSegments.Count && rightCaseSensitiveSegments[i]); + var left = preserveCase ? leftSegments[i] : NameFold.Fold(leftSegments[i]) ?? leftSegments[i]; + var right = preserveCase ? rightSegments[i] : NameFold.Fold(rightSegments[i]) ?? rightSegments[i]; + if (!string.Equals(left, right, StringComparison.Ordinal)) + return false; + } + + return true; + } + + private static bool SegmentsEqual(string left, bool leftCaseSensitive, string right, bool rightCaseSensitive) + => string.Equals(left, right, leftCaseSensitive || rightCaseSensitive ? StringComparison.Ordinal : StringComparison.OrdinalIgnoreCase); } diff --git a/tests/CodeIndex.Tests/ReferenceExtractorTests.cs b/tests/CodeIndex.Tests/ReferenceExtractorTests.cs index 9614904ef6..3a1a111779 100644 --- a/tests/CodeIndex.Tests/ReferenceExtractorTests.cs +++ b/tests/CodeIndex.Tests/ReferenceExtractorTests.cs @@ -25088,6 +25088,35 @@ public void Extract_SqlCallBacktickReservedWord_SurvivesQuoteStripping() Assert.Contains(references, r => r.SymbolName == "select" && r.ReferenceKind == "call" && r.Line == 2); } + [Fact] + public void SqlNameResolver_QuotedQualifiedNames_PreserveDialectSpecificMatching() + { + const string mysqlContext = "SELECT * FROM `mydb`.`mytbl`;"; + const int mysqlColumn = 28; + Assert.Equal("mydb.mytbl", SqlNameResolver.ResolveReferenceNameAtColumn("mytbl", mysqlContext, null, mysqlColumn)); + Assert.True(SqlNameResolver.ContextContainsQualifiedNameFoldedAtColumn(mysqlContext, "`MYDB`.`MYTBL`", mysqlColumn)); + + const string tsqlContext = "SELECT * FROM [sales data].[order table];"; + const int tsqlColumn = 35; + Assert.Equal("sales data.order table", SqlNameResolver.ResolveReferenceNameAtColumn("order table", tsqlContext, null, tsqlColumn)); + Assert.True(SqlNameResolver.ContextContainsQualifiedNameFoldedAtColumn(tsqlContext, "[SALES DATA].[ORDER TABLE]", tsqlColumn)); + + const string postgresContext = "SELECT * FROM \"Sales\".\"Orders\";"; + const int postgresColumn = 24; + Assert.Equal("Sales.Orders", SqlNameResolver.ResolveReferenceNameAtColumn("Orders", postgresContext, null, postgresColumn)); + Assert.True(SqlNameResolver.ContextContainsQualifiedNameFoldedAtColumn(postgresContext, "\"Sales\".\"Orders\"", postgresColumn)); + Assert.False(SqlNameResolver.ContextContainsQualifiedNameFoldedAtColumn(postgresContext, "\"sales\".\"orders\"", postgresColumn)); + Assert.True(SqlNameResolver.ReferenceMatchesTargetAtColumn("Orders", postgresContext, null, postgresColumn, "Sales.Orders")); + Assert.False(SqlNameResolver.ReferenceMatchesTargetAtColumn("Orders", postgresContext, null, postgresColumn, "sales.orders")); + + const string mixedPostgresContext = "SELECT * FROM \"Sales\".orders;"; + const int mixedPostgresColumn = 24; + Assert.True(SqlNameResolver.ContextContainsQualifiedNameFoldedAtColumn(mixedPostgresContext, "\"Sales\".ORDERS", mixedPostgresColumn)); + Assert.False(SqlNameResolver.ContextContainsQualifiedNameFoldedAtColumn(mixedPostgresContext, "\"sales\".orders", mixedPostgresColumn)); + Assert.True(SqlNameResolver.ReferenceMatchesTargetAtColumn("orders", mixedPostgresContext, null, mixedPostgresColumn, "\"Sales\".ORDERS")); + Assert.False(SqlNameResolver.ReferenceMatchesTargetAtColumn("orders", mixedPostgresContext, null, mixedPostgresColumn, "\"sales\".orders")); + } + [Fact] public void Extract_SqlHashCommentedCall_DoesNotEmitReference() {