Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@ Flux<ChatResponse> fineSelect(SchemaDTO schemaDTO, String query, String evidence
String sqlGenerateSchemaMissingAdvice, DbConfigBO specificDbConfig, Consumer<SchemaDTO> dtoConsumer);

default String sqlTrim(String sql) {
return MarkdownParserUtil.extractRawText(sql).trim();
return MarkdownParserUtil.extractLastRawText(sql).trim();
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -25,15 +25,34 @@ public static String extractText(String markdownCode) {
}

public static String extractRawText(String markdownCode) {
// Find the start of a code block (3 or more backticks)
CodeBlock codeBlock = findNextCodeBlock(markdownCode, 0);
return codeBlock == null ? markdownCode : codeBlock.rawText;
}

public static String extractLastRawText(String markdownCode) {
CodeBlock lastCodeBlock = null;
int searchIndex = 0;

while (searchIndex <= markdownCode.length() - 3) {
CodeBlock codeBlock = findNextCodeBlock(markdownCode, searchIndex);
if (codeBlock == null) {
break;
}
lastCodeBlock = codeBlock;
searchIndex = codeBlock.nextSearchIndex;
}

return lastCodeBlock == null ? markdownCode : lastCodeBlock.rawText;
}

private static CodeBlock findNextCodeBlock(String markdownCode, int searchIndex) {
int startIndex = -1;
int delimiterLength = 0;

for (int i = 0; i <= markdownCode.length() - 3; i++) {
for (int i = searchIndex; i <= markdownCode.length() - 3; i++) {
if (markdownCode.substring(i, i + 3).equals("```")) {
startIndex = i;
delimiterLength = 3;
// Count additional backticks
while (i + delimiterLength < markdownCode.length() && markdownCode.charAt(i + delimiterLength) == '`') {
delimiterLength++;
}
Expand All @@ -42,29 +61,38 @@ public static String extractRawText(String markdownCode) {
}

if (startIndex == -1) {
return markdownCode; // No code block found
return null;
}

// Skip the opening delimiter and optional language specification
int contentStart = startIndex + delimiterLength;
while (contentStart < markdownCode.length() && markdownCode.charAt(contentStart) != '\n') {
contentStart++;
}
if (contentStart < markdownCode.length() && markdownCode.charAt(contentStart) == '\n') {
contentStart++; // Skip the newline after language spec
contentStart++;
}

// Find the closing delimiter
String closingDelimiter = "`".repeat(delimiterLength);
int endIndex = markdownCode.indexOf(closingDelimiter, contentStart);

if (endIndex == -1) {
// No closing delimiter found, return from content start to end
return markdownCode.substring(contentStart);
return new CodeBlock(markdownCode.substring(contentStart), markdownCode.length());
}

return new CodeBlock(markdownCode.substring(contentStart, endIndex), endIndex + delimiterLength);
}

private static final class CodeBlock {

private final String rawText;

private final int nextSearchIndex;

private CodeBlock(String rawText, int nextSearchIndex) {
this.rawText = rawText;
this.nextSearchIndex = nextSearchIndex;
}

// Extract just the content between delimiters
return markdownCode.substring(contentStart, endIndex);
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -202,10 +202,46 @@ void sqlTrim_markdownCodeBlockWithLineBreaks_returnsCleanSql() {
}

@Test
void sqlTrim_multipleBacktickBlocks_extractsFirst() {
String sql = "```sql\nSELECT 1\n```\nSome text\n```sql\nSELECT 2\n```";
void sqlTrim_multipleSqlBlocks_returnsLastSqlBlock() {
String sql = """
```sql
SELECT create_by
FROM sys_dept
WHERE dept_name = '研发部'
```
Some text
```sql
SELECT user_id
FROM sys_user
WHERE user_id IN (
SELECT create_by
FROM sys_dept
WHERE dept_name = '研发部'
)
```
More text
```sql
SELECT *
FROM sys_user
WHERE user_id IN (
SELECT create_by
FROM sys_dept
WHERE dept_name = '研发部'
)
```
""";
String expected = """
SELECT *
FROM sys_user
WHERE user_id IN (
SELECT create_by
FROM sys_dept
WHERE dept_name = '研发部'
)
""".trim();

String result = nl2SqlService.sqlTrim(sql);
assertNotNull(result);
assertEquals(expected, result);
}

@Test
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,24 @@ void testExtractRawText_moreBackticks() {
assertEquals("code here\n", MarkdownParserUtil.extractRawText(input));
}

@Test
void testExtractLastRawText_singleBlock() {
String input = "```sql\nSELECT 1\n```";
assertEquals("SELECT 1\n", MarkdownParserUtil.extractLastRawText(input));
}

@Test
void testExtractLastRawText_multipleBlocks_returnsLast() {
String input = "```sql\nSELECT 1\n```\n```sql\nSELECT 2\n```";
assertEquals("SELECT 2\n", MarkdownParserUtil.extractLastRawText(input));
}

@Test
void testExtractLastRawText_noCodeBlock_returnsOriginal() {
String input = "plain SQL text";
assertEquals("plain SQL text", MarkdownParserUtil.extractLastRawText(input));
}

@Test
void testExtractText_replacesNewlines() {
String input = "```sql\nSELECT *\nFROM users;\n```";
Expand Down
Loading