feat(completion): add scoped metadata assistant

This commit is contained in:
t8y2 2026-06-24 00:38:18 +08:00
parent 3f65297f25
commit cebb39cb80
37 changed files with 2709 additions and 85 deletions

View File

@ -60,6 +60,18 @@ Requires JDK 8 and 21 (Gradle toolchain auto-downloads if needed).
Output JARs are in `drivers/{module}/build/libs/`. Native agents build from `drivers/oracle-go` and `drivers/xugu`.
### Local DBX Runtime Test
When changing `agents/drivers/<db_type>/` or shared Java agent protocol code, rebuild the target agent and replace the runtime JAR used by the local DBX app:
```bash
./gradlew :<db_type>:shadowJar
cp ~/.dbx/agents/drivers/<db_type>/agent.jar ~/.dbx/agents/drivers/<db_type>/agent.jar.bak
cp agents/drivers/<db_type>/build/libs/*-all.jar ~/.dbx/agents/drivers/<db_type>/agent.jar
```
Restart DBX or disconnect and reconnect the database so the new agent process loads the replacement JAR.
## Development
- Agent authoring guide: [docs/agent-authoring.md](docs/agent-authoring.md)

View File

@ -15,6 +15,7 @@ public final class AgentProtocol {
public static final String METHOD_LIST_SCHEMAS = "list_schemas";
public static final String METHOD_LIST_TABLES = "list_tables";
public static final String METHOD_LIST_OBJECTS = "list_objects";
public static final String METHOD_COMPLETION_ASSISTANT_SEARCH_V1 = "completion_assistant_search_v1";
public static final String METHOD_GET_OBJECT_SOURCE = "get_object_source";
public static final String METHOD_GET_TABLE_DDL = "get_table_ddl";
public static final String METHOD_GET_COLUMNS = "get_columns";
@ -81,6 +82,7 @@ public final class AgentProtocol {
METHOD_LIST_SCHEMAS,
METHOD_LIST_TABLES,
METHOD_LIST_OBJECTS,
METHOD_COMPLETION_ASSISTANT_SEARCH_V1,
METHOD_GET_OBJECT_SOURCE,
METHOD_GET_TABLE_DDL,
METHOD_GET_COLUMNS,

View File

@ -20,6 +20,11 @@ public abstract class BaseDatabaseAgent implements DatabaseAgent {
throw new UnsupportedOperationException("Object source is not supported");
}
@Override
public CompletionAssistantResponse completionAssistantSearch(CompletionAssistantRequest request) {
throw new UnsupportedOperationException("Completion assistant search is not supported by this agent");
}
@Override
public String getTableDdl(String schema, String table) {
List<IndexInfo> indexes;

View File

@ -0,0 +1,41 @@
package com.dbx.agent;
public final class CompletionAssistantCandidate {
private final String name;
private final CompletionAssistantCandidateKind kind;
private final String database;
private final String schema;
private final String parent_schema;
private final String parent_name;
private final String comment;
private final String data_type;
public CompletionAssistantCandidate(
String name,
CompletionAssistantCandidateKind kind,
String database,
String schema,
String parentSchema,
String parentName,
String comment,
String dataType
) {
this.name = name;
this.kind = kind;
this.database = database;
this.schema = schema;
this.parent_schema = parentSchema;
this.parent_name = parentName;
this.comment = comment;
this.data_type = dataType;
}
public String getName() { return name; }
public CompletionAssistantCandidateKind getKind() { return kind; }
public String getDatabase() { return database; }
public String getSchema() { return schema; }
public String getParent_schema() { return parent_schema; }
public String getParent_name() { return parent_name; }
public String getComment() { return comment; }
public String getData_type() { return data_type; }
}

View File

@ -0,0 +1,22 @@
package com.dbx.agent;
import com.google.gson.annotations.SerializedName;
public enum CompletionAssistantCandidateKind {
@SerializedName("database")
DATABASE,
@SerializedName("schema")
SCHEMA,
@SerializedName("table")
TABLE,
@SerializedName("view")
VIEW,
@SerializedName("procedure")
PROCEDURE,
@SerializedName("function")
FUNCTION,
@SerializedName("column")
COLUMN,
@SerializedName("object")
OBJECT
}

View File

@ -0,0 +1,10 @@
package com.dbx.agent;
import com.google.gson.annotations.SerializedName;
public enum CompletionAssistantMatchMode {
@SerializedName("prefix")
PREFIX,
@SerializedName("contains")
CONTAINS
}

View File

@ -0,0 +1,22 @@
package com.dbx.agent;
import com.google.gson.annotations.SerializedName;
public enum CompletionAssistantObjectKind {
@SerializedName("database")
DATABASE,
@SerializedName("schema")
SCHEMA,
@SerializedName("table")
TABLE,
@SerializedName("view")
VIEW,
@SerializedName("routine")
ROUTINE,
@SerializedName("procedure")
PROCEDURE,
@SerializedName("function")
FUNCTION,
@SerializedName("column")
COLUMN
}

View File

@ -0,0 +1,34 @@
package com.dbx.agent;
import java.util.ArrayList;
import java.util.List;
public final class CompletionAssistantRequest {
private String connection_id;
private String database;
private String schema;
private List<CompletionAssistantObjectKind> object_kinds = new ArrayList<>();
private String mask = "";
private boolean case_sensitive;
private boolean global_search;
private Integer max_results;
private boolean search_in_comments;
private boolean search_in_definitions;
private String parent_schema;
private String parent_name;
private CompletionAssistantMatchMode match_mode;
public String getConnection_id() { return connection_id; }
public String getDatabase() { return database == null ? "" : database; }
public String getSchema() { return schema; }
public List<CompletionAssistantObjectKind> getObject_kinds() { return object_kinds == null ? new ArrayList<>() : object_kinds; }
public String getMask() { return mask == null ? "" : mask; }
public boolean getCase_sensitive() { return case_sensitive; }
public boolean getGlobal_search() { return global_search; }
public Integer getMax_results() { return max_results; }
public boolean getSearch_in_comments() { return search_in_comments; }
public boolean getSearch_in_definitions() { return search_in_definitions; }
public String getParent_schema() { return parent_schema; }
public String getParent_name() { return parent_name; }
public CompletionAssistantMatchMode getMatch_mode() { return match_mode; }
}

View File

@ -0,0 +1,19 @@
package com.dbx.agent;
import java.util.List;
public final class CompletionAssistantResponse {
private final List<CompletionAssistantCandidate> candidates;
private final boolean incomplete;
private final boolean fallback_used;
public CompletionAssistantResponse(List<CompletionAssistantCandidate> candidates, boolean incomplete, boolean fallbackUsed) {
this.candidates = candidates;
this.incomplete = incomplete;
this.fallback_used = fallbackUsed;
}
public List<CompletionAssistantCandidate> getCandidates() { return candidates; }
public boolean getIncomplete() { return incomplete; }
public boolean getFallback_used() { return fallback_used; }
}

View File

@ -60,6 +60,11 @@ public abstract class ConfiguredJdbcAgent extends AbstractJdbcAgent {
return StandardJdbcMetadata.INSTANCE.listObjects(listTables(schema), schema);
}
@Override
public CompletionAssistantResponse completionAssistantSearch(CompletionAssistantRequest request) {
return StandardJdbcMetadata.INSTANCE.completionAssistantSearch(requireConnection(), profile, configuredDatabase, request);
}
@Override
public ObjectSource getObjectSource(String schema, String name, String objectType) {
throw new UnsupportedOperationException("Object source is not supported");

View File

@ -24,6 +24,10 @@ public interface DatabaseAgent {
return result;
}
default CompletionAssistantResponse completionAssistantSearch(CompletionAssistantRequest request) {
throw new UnsupportedOperationException("Completion assistant search is not supported by this agent");
}
List<ColumnInfo> getColumns(String schema, String table);
default ObjectSource getObjectSource(String schema, String name, String objectType) {

View File

@ -115,6 +115,10 @@ public final class JsonRpcServer {
switchCatalog(params);
return agent.listObjects(params.get("schema").getAsString());
}
if (AgentProtocol.METHOD_COMPLETION_ASSISTANT_SEARCH_V1.equals(method)) {
switchCatalog(params);
return agent.completionAssistantSearch(gson.fromJson(params, CompletionAssistantRequest.class));
}
if (AgentProtocol.METHOD_GET_OBJECT_SOURCE.equals(method)) {
switchCatalog(params);
return agent.getObjectSource(

View File

@ -90,6 +90,40 @@ public final class StandardJdbcMetadata {
return result;
}
public CompletionAssistantResponse completionAssistantSearch(
Connection conn,
JdbcAgentProfile profile,
String configuredDatabase,
CompletionAssistantRequest request
) {
return unchecked(() -> {
int limit = boundedLimit(request.getMax_results());
List<CompletionAssistantCandidate> candidates = new ArrayList<>();
List<CompletionAssistantObjectKind> kinds = normalizedObjectKinds(request);
DatabaseMetaData meta = conn.getMetaData();
String catalog = blankToNull(request.getDatabase());
String schema = completionSchema(request);
if (containsKind(kinds, CompletionAssistantObjectKind.SCHEMA)) {
appendCompletionSchemas(candidates, meta, profile, request, limit);
}
if (candidates.size() < limit && containsTableLike(kinds)) {
appendCompletionTables(candidates, meta, profile, catalog, schema, request, kinds, limit);
if (candidates.isEmpty() && profile.getCatalogFallbackEnabled() && configuredDatabase != null && !configuredDatabase.trim().isEmpty()) {
appendCompletionTables(candidates, meta, profile, configuredDatabase, schema, request, kinds, limit);
}
}
if (candidates.size() < limit && containsKind(kinds, CompletionAssistantObjectKind.COLUMN)) {
appendCompletionColumns(candidates, meta, catalog, schema, request, limit);
if (candidates.isEmpty() && profile.getCatalogFallbackEnabled() && configuredDatabase != null && !configuredDatabase.trim().isEmpty()) {
appendCompletionColumns(candidates, meta, configuredDatabase, schema, request, limit);
}
}
return new CompletionAssistantResponse(candidates, candidates.size() >= limit, false);
});
}
public List<ColumnInfo> getColumns(Connection conn, JdbcAgentProfile profile, String configuredDatabase, String schema, String table) {
return unchecked(() -> {
DatabaseMetaData meta = conn.getMetaData();
@ -229,6 +263,188 @@ public final class StandardJdbcMetadata {
}
}
private void appendCompletionSchemas(
List<CompletionAssistantCandidate> result,
DatabaseMetaData meta,
JdbcAgentProfile profile,
CompletionAssistantRequest request,
int limit
) throws Exception {
Set<String> names = new LinkedHashSet<>();
try {
appendSchemas(names, meta.getSchemas(null, null));
} catch (Exception | AbstractMethodError first) {
try {
appendSchemas(names, meta.getSchemas());
} catch (Exception | AbstractMethodError ignored) {
}
}
List<String> sorted = new ArrayList<>(names);
Collections.sort(sorted);
for (String name : sorted) {
if (result.size() >= limit) {
return;
}
if (profile.getExcludedSchemas().contains(name.toUpperCase(Locale.ROOT)) || !completionNameMatches(name, request)) {
continue;
}
result.add(new CompletionAssistantCandidate(
name,
CompletionAssistantCandidateKind.SCHEMA,
request.getDatabase(),
name,
null,
null,
null,
null
));
}
}
private void appendCompletionTables(
List<CompletionAssistantCandidate> result,
DatabaseMetaData meta,
JdbcAgentProfile profile,
String catalog,
String schema,
CompletionAssistantRequest request,
List<CompletionAssistantObjectKind> kinds,
int limit
) throws Exception {
String[] tableTypes = getDriverTableTypes(meta, profile);
try (ResultSet rs = meta.getTables(catalog, blankToNull(schema), completionPattern(request), tableTypes)) {
while (rs.next() && result.size() < limit) {
String name = rs.getString("TABLE_NAME");
String type = normalizeTableType(rs.getString("TABLE_TYPE"));
CompletionAssistantCandidateKind kind = completionTableKind(type);
if (!completionTableKindAllowed(kind, kinds) || !completionNameMatches(name, request)) {
continue;
}
result.add(new CompletionAssistantCandidate(
name,
kind,
request.getDatabase(),
schema,
null,
null,
rs.getString("REMARKS"),
null
));
}
}
result.sort(Comparator.comparing(CompletionAssistantCandidate::getName));
}
private void appendCompletionColumns(
List<CompletionAssistantCandidate> result,
DatabaseMetaData meta,
String catalog,
String schema,
CompletionAssistantRequest request,
int limit
) throws Exception {
String table = request.getParent_name();
if (table == null || table.trim().isEmpty()) {
return;
}
try (ResultSet rs = meta.getColumns(catalog, blankToNull(schema), table, completionPattern(request))) {
while (rs.next() && result.size() < limit) {
String name = rs.getString("COLUMN_NAME");
if (!completionNameMatches(name, request)) {
continue;
}
result.add(new CompletionAssistantCandidate(
name,
CompletionAssistantCandidateKind.COLUMN,
request.getDatabase(),
schema,
schema,
table,
rs.getString("REMARKS"),
rs.getString("TYPE_NAME")
));
}
}
}
private static List<CompletionAssistantObjectKind> normalizedObjectKinds(CompletionAssistantRequest request) {
List<CompletionAssistantObjectKind> kinds = request.getObject_kinds();
if (kinds.isEmpty()) {
kinds = new ArrayList<>();
kinds.add(CompletionAssistantObjectKind.TABLE);
kinds.add(CompletionAssistantObjectKind.VIEW);
}
return kinds;
}
private static String completionSchema(CompletionAssistantRequest request) {
String parentSchema = request.getParent_schema();
if (parentSchema != null && !parentSchema.trim().isEmpty()) {
return parentSchema;
}
return request.getSchema();
}
private static String completionPattern(CompletionAssistantRequest request) {
String mask = request.getMask();
if (mask.trim().isEmpty()) {
return "%";
}
String escaped = mask.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_");
if (request.getMatch_mode() == CompletionAssistantMatchMode.CONTAINS) {
return "%" + escaped + "%";
}
return escaped + "%";
}
private static boolean completionNameMatches(String name, CompletionAssistantRequest request) {
if (name == null) {
return false;
}
String mask = request.getMask();
if (mask.trim().isEmpty()) {
return true;
}
String candidate = request.getCase_sensitive() ? name : name.toLowerCase(Locale.ROOT);
String expected = request.getCase_sensitive() ? mask : mask.toLowerCase(Locale.ROOT);
if (request.getMatch_mode() == CompletionAssistantMatchMode.CONTAINS) {
return candidate.contains(expected);
}
return candidate.startsWith(expected);
}
private static int boundedLimit(Integer requested) {
if (requested == null) {
return 100;
}
return Math.max(1, Math.min(1000, requested));
}
private static boolean containsKind(List<CompletionAssistantObjectKind> kinds, CompletionAssistantObjectKind kind) {
return kinds.contains(kind);
}
private static boolean containsTableLike(List<CompletionAssistantObjectKind> kinds) {
return kinds.contains(CompletionAssistantObjectKind.TABLE) || kinds.contains(CompletionAssistantObjectKind.VIEW);
}
private static CompletionAssistantCandidateKind completionTableKind(String type) {
if ("VIEW".equalsIgnoreCase(type)) {
return CompletionAssistantCandidateKind.VIEW;
}
return CompletionAssistantCandidateKind.TABLE;
}
private static boolean completionTableKindAllowed(
CompletionAssistantCandidateKind kind,
List<CompletionAssistantObjectKind> requestedKinds
) {
if (kind == CompletionAssistantCandidateKind.VIEW) {
return requestedKinds.contains(CompletionAssistantObjectKind.VIEW);
}
return requestedKinds.contains(CompletionAssistantObjectKind.TABLE);
}
private Set<String> primaryKeys(DatabaseMetaData meta, String catalog, String schema, String table) {
Set<String> keys = new LinkedHashSet<>();
try (ResultSet rs = meta.getPrimaryKeys(catalog, blankToNull(schema), table)) {

View File

@ -43,6 +43,7 @@
"list_schemas",
"list_tables",
"list_objects",
"completion_assistant_search_v1",
"get_object_source",
"get_table_ddl",
"get_columns",

View File

@ -247,6 +247,44 @@ class StandardJdbcMetadataTest {
assertEquals(Collections.emptyList(), StandardJdbcMetadata.INSTANCE.listTriggers("APP", "ORDERS"));
}
@Test
void completionAssistantSearchesTablesAndColumnsWithServerSideMasks() {
AtomicReference<Object[]> capturedTableArgs = new AtomicReference<>();
AtomicReference<Object[]> capturedColumnArgs = new AtomicReference<>();
Connection conn = connection(
rows(),
rows(
row("TABLE_NAME", "ACCOUNTS", "TABLE_TYPE", "TABLE", "REMARKS", "account table"),
row("TABLE_NAME", "ACCOUNT_VIEW", "TABLE_TYPE", "VIEW", "REMARKS", null)
),
rows(),
rows(row("COLUMN_NAME", "DISPLAY_NAME", "TYPE_NAME", "VARCHAR", "NULLABLE", DatabaseMetaData.columnNullable, "COLUMN_DEF", null, "REMARKS", "display")),
rows(),
rows(),
UnsupportedSchemaCall.NONE,
rows(row("TABLE_TYPE", "TABLE"), row("TABLE_TYPE", "VIEW")),
null,
capturedTableArgs,
capturedColumnArgs
);
CompletionAssistantRequest tablesRequest = request("sales", "APP", "ACC", Arrays.asList(CompletionAssistantObjectKind.TABLE, CompletionAssistantObjectKind.VIEW), null);
CompletionAssistantResponse tables = StandardJdbcMetadata.INSTANCE.completionAssistantSearch(conn, profile, "sales", tablesRequest);
assertEquals(2, tables.getCandidates().size());
assertEquals(CompletionAssistantCandidateKind.TABLE, tables.getCandidates().get(0).getKind());
assertEquals("ACC%", capturedTableArgs.get()[2]);
CompletionAssistantRequest columnsRequest = request("sales", "APP", "DISPLAY", Collections.singletonList(CompletionAssistantObjectKind.COLUMN), "ACCOUNTS");
CompletionAssistantResponse columns = StandardJdbcMetadata.INSTANCE.completionAssistantSearch(conn, profile, "sales", columnsRequest);
assertEquals(1, columns.getCandidates().size());
assertEquals("DISPLAY_NAME", columns.getCandidates().get(0).getName());
assertEquals("VARCHAR", columns.getCandidates().get(0).getData_type());
assertEquals("ACCOUNTS", capturedColumnArgs.get()[2]);
assertEquals("DISPLAY%", capturedColumnArgs.get()[3]);
}
private static Connection connection(
ResultSet schemas,
ResultSet tables,
@ -315,6 +353,22 @@ class StandardJdbcMetadataTest {
UnsupportedSchemaCall unsupportedSchemaCall,
ResultSet tableTypes,
AtomicReference<String[]> capturedTableTypes
) {
return connection(schemas, tables, primaryKeys, columns, indexes, foreignKeys, unsupportedSchemaCall, tableTypes, capturedTableTypes, null, null);
}
private static Connection connection(
ResultSet schemas,
ResultSet tables,
ResultSet primaryKeys,
ResultSet columns,
ResultSet indexes,
ResultSet foreignKeys,
UnsupportedSchemaCall unsupportedSchemaCall,
ResultSet tableTypes,
AtomicReference<String[]> capturedTableTypes,
AtomicReference<Object[]> capturedTableArgs,
AtomicReference<Object[]> capturedColumnArgs
) {
DatabaseMetaData meta = proxy(DatabaseMetaData.class, new MethodHandler() {
@Override
@ -327,6 +381,9 @@ class StandardJdbcMetadataTest {
return schemas;
}
if ("getTables".equals(name)) {
if (capturedTableArgs != null) {
capturedTableArgs.set(args);
}
if (capturedTableTypes != null && args != null && args.length > 3) {
capturedTableTypes.set((String[]) args[3]);
}
@ -339,6 +396,9 @@ class StandardJdbcMetadataTest {
return primaryKeys;
}
if ("getColumns".equals(name)) {
if (capturedColumnArgs != null) {
capturedColumnArgs.set(args);
}
return columns;
}
if ("getIndexInfo".equals(name)) {
@ -420,6 +480,35 @@ class StandardJdbcMetadataTest {
return row;
}
private static CompletionAssistantRequest request(
String database,
String schema,
String mask,
List<CompletionAssistantObjectKind> kinds,
String parentName
) {
CompletionAssistantRequest request = new CompletionAssistantRequest();
setField(request, "database", database);
setField(request, "schema", schema);
setField(request, "mask", mask);
setField(request, "object_kinds", kinds);
setField(request, "parent_schema", schema);
setField(request, "parent_name", parentName);
setField(request, "max_results", 10);
setField(request, "match_mode", CompletionAssistantMatchMode.PREFIX);
return request;
}
private static void setField(Object target, String name, Object value) {
try {
java.lang.reflect.Field field = target.getClass().getDeclaredField(name);
field.setAccessible(true);
field.set(target, value);
} catch (Exception e) {
throw new RuntimeException(e);
}
}
private static <T> T proxy(Class<T> type, final MethodHandler handler) {
InvocationHandler invocationHandler = new InvocationHandler() {
@Override

View File

@ -46,7 +46,7 @@ import * as api from "@/lib/api";
import { areSqlSemanticDiagnosticsEqual, buildSqlParserErrorDiagnostic, buildSqlSemanticDiagnostics, shouldRunSqlSemanticDiagnostics, type SqlSemanticDiagnostic } from "@/lib/sqlSemanticDiagnostics";
import { buildRedisSyntaxDiagnostics, shouldRunRedisDiagnostics } from "@/lib/redisSyntaxDiagnostics";
import { buildRedisCompletionItemsFromContext, getRedisCompletionContext, getRedisCompletionResultValidFor, shouldAutoOpenRedisCompletion, takesKeyArgument, type RedisCompletionItem } from "@/lib/redisCompletion";
import type { SqlCompletionColumn, SqlCompletionForeignKey, SqlCompletionItem, SqlCompletionObject } from "@/lib/sqlCompletion";
import type { SqlCompletionColumn, SqlCompletionForeignKey, SqlCompletionItem, SqlCompletionObject, SqlCompletionTable } from "@/lib/sqlCompletion";
import type { DatabaseType, SqlReferenceAnalysis, SqlTableReference, SqlTextSpan } from "@/types/database";
const props = defineProps<{
@ -65,6 +65,8 @@ const props = defineProps<{
initialSelection?: { anchor: number; head: number };
}>();
const COMPLETION_REMOTE_LATENCY_BUDGET_MS = 120;
const emit = defineEmits<{
"update:modelValue": [value: string];
selectionChange: [value: string];
@ -1490,6 +1492,20 @@ function mergeCompletionTables(existing: Array<{ name: string; schema?: string;
return merged;
}
function withCompletionLatencyBudget<T>(remote: Promise<T>, local: T): Promise<T> {
return Promise.race([remote, new Promise<T>((resolve) => setTimeout(() => resolve(local), COMPLETION_REMOTE_LATENCY_BUDGET_MS))]);
}
function listCompletionTablesWithLatencyBudget(connectionId: string, database: string, filter: string, limit: number, schema?: string): Promise<SqlCompletionTable[]> {
const local = connectionStore.lookupLocalCompletionTables(connectionId, database, filter, limit, schema);
const remote = connectionStore.listCompletionTables(connectionId, database, filter, limit, schema).then((tables) => {
cachedTables = mergeCompletionTables(cachedTables, tables);
return tables;
});
if (local.length === 0) return remote;
return withCompletionLatencyBudget(remote, local);
}
async function performAsyncCompletionWithResult(epoch: number, completionContext: ReturnType<typeof getSqlCompletionContext>, fullDoc: string, position: number) {
const localOnlyMetadata = usesLocalOnlyCompletionMetadata();
const onDemandOnlyColumns = usesOnDemandOnlyCompletionColumns();
@ -1526,7 +1542,7 @@ async function performAsyncCompletionWithResult(epoch: number, completionContext
let tables = shouldLoadTables
? localOnlyMetadata
? connectionStore.lookupLocalCompletionTables(props.connectionId!, tableLookupDatabase, tableLookupFilter, MAX_COMPLETION_TABLES, tableLookupSchema)
: await connectionStore.listCompletionTables(props.connectionId!, tableLookupDatabase, tableLookupFilter, MAX_COMPLETION_TABLES, tableLookupSchema)
: await listCompletionTablesWithLatencyBudget(props.connectionId!, tableLookupDatabase, tableLookupFilter, MAX_COMPLETION_TABLES, tableLookupSchema)
: cachedTables;
if (epoch !== completionEpoch) return null;
@ -1568,7 +1584,7 @@ async function performAsyncCompletionWithResult(epoch: number, completionContext
if (completionContext.qualifier && !qualifierDatabase && !isReferencedTableQualifier(completionContext) && tables.length === 0 && (completionContext.suggestTables || completionContext.exclusiveColumnSuggestions)) {
const schemaTables = localOnlyMetadata
? connectionStore.lookupLocalCompletionTables(props.connectionId!, props.database!, completionContext.prefix, MAX_COMPLETION_TABLES, completionContext.qualifier)
: await connectionStore.listCompletionTables(props.connectionId!, props.database!, completionContext.prefix, MAX_COMPLETION_TABLES, completionContext.qualifier);
: await listCompletionTablesWithLatencyBudget(props.connectionId!, props.database!, completionContext.prefix, MAX_COMPLETION_TABLES, completionContext.qualifier);
if (schemaTables.length > 0) {
tables = schemaTables;
qualifierIsSchema = true;

View File

@ -27,3 +27,100 @@ describe("sqlCompletion quoted schema qualifiers", () => {
expect(items.some((item) => item.label === "shipments" && item.type === "table")).toBe(true);
});
});
describe("sqlCompletion scoped context classification", () => {
it("classifies JOIN table contexts", () => {
const sql = "SELECT * FROM users u JOIN ";
const context = getSqlCompletionContext(sql, sql.length);
expect(context.contextKind).toBe("join");
expect(context.suggestTables).toBe(true);
expect(context.exclusiveTableSuggestions).toBe(true);
});
it("classifies alias-qualified column contexts", () => {
const sql = "SELECT * FROM users u WHERE u.";
const context = getSqlCompletionContext(sql, sql.length);
expect(context.contextKind).toBe("alias_column");
expect(context.qualifier).toBe("u");
expect(context.suggestColumns).toBe(true);
});
it("classifies CALL routine contexts", () => {
const sql = "CALL usp_";
const context = getSqlCompletionContext(sql, sql.length);
expect(context.contextKind).toBe("exec");
expect(context.suggestRoutines).toBe(true);
expect(context.exclusiveRoutineSuggestions).toBe(true);
});
it("classifies INSERT column-list contexts", () => {
const sql = "INSERT INTO dbo.Users (";
const context = getSqlCompletionContext(sql, sql.length);
expect(context.contextKind).toBe("column");
expect(context.insertSchema).toBe("dbo");
expect(context.insertTable).toBe("users");
expect(context.exclusiveColumnSuggestions).toBe(true);
});
it("classifies UPDATE SET column contexts", () => {
const sql = "UPDATE dbo.Users SET ";
const context = getSqlCompletionContext(sql, sql.length);
expect(context.contextKind).toBe("column");
expect(context.updateTarget).toEqual({ schema: "dbo", table: "Users" });
expect(context.suggestColumns).toBe(true);
});
it("extracts statement-local table aliases", () => {
const sql = "SELECT * FROM dbo.Users u JOIN Orders AS o ON o.user_id = u.id WHERE u.";
const context = getSqlCompletionContext(sql, sql.length);
expect(context.referencedTables).toEqual(expect.arrayContaining([expect.objectContaining({ schema: "dbo", name: "Users", alias: "u" }), expect.objectContaining({ name: "Orders", alias: "o" })]));
});
it("exposes CTEs as table-like referenced tables", () => {
const sql = "WITH recent_orders(id, total) AS (SELECT id, total FROM orders) SELECT * FROM recent_orders ro WHERE ro.";
const context = getSqlCompletionContext(sql, sql.length);
expect(context.referencedTables).toEqual(expect.arrayContaining([expect.objectContaining({ name: "recent_orders", columns: ["id", "total"] }), expect.objectContaining({ name: "recent_orders", alias: "ro" })]));
});
it("extracts subquery aliases and projected columns", () => {
const sql = "SELECT * FROM (SELECT id, name AS user_name FROM users) sq WHERE sq.";
const context = getSqlCompletionContext(sql, sql.length);
expect(context.referencedTables).toEqual(expect.arrayContaining([expect.objectContaining({ name: "sq", alias: "sq", columns: ["id", "user_name"] })]));
});
});
describe("sqlCompletion scoped metadata ranking", () => {
it("ranks exact and prefix table matches ahead of contains/fuzzy matches", () => {
const sql = "SELECT * FROM Temp";
const items = buildSqlCompletionItems(sql, sql.length, {
dialect: "sqlserver",
tables: [
{ name: "ArchiveTempTable", schema: "dbo", type: "table" },
{ name: "TempAudit", schema: "dbo", type: "table" },
{ name: "Temp", schema: "dbo", type: "table" },
{ name: "Template", schema: "dbo", type: "table" },
],
columnsByTable: new Map(),
}).filter((item) => item.type === "table");
expect(items.map((item) => item.label).slice(0, 3)).toEqual(["Temp", "Template", "TempAudit"]);
expect(items.some((item) => item.label === "ArchiveTempTable")).toBe(true);
});
it("keeps large table catalogs bounded", () => {
const tables = Array.from({ length: 500 }, (_, index) => ({ name: `TempTable_${String(index).padStart(3, "0")}`, schema: "dbo", type: "table" as const }));
const sql = "SELECT * FROM Temp";
const items = buildSqlCompletionItems(sql, sql.length, { dialect: "sqlserver", tables, columnsByTable: new Map() }).filter((item) => item.type === "table");
expect(items.length).toBeLessThanOrEqual(200);
expect(items[0]?.label).toBe("TempTable_000");
});
});

View File

@ -118,6 +118,7 @@ export const getTableComment = forward("getTableComment");
export const listObjects = forward("listObjects");
export const listObjectStatistics = forward("listObjectStatistics");
export const listCompletionObjects = forward("listCompletionObjects");
export const completionAssistantSearch = forward("completionAssistantSearch");
export const getObjectSource = forward("getObjectSource");
export const getColumns = forward("getColumns");
export const listIndexes = forward("listIndexes");

View File

@ -4,6 +4,8 @@ import type {
LinkedServerInfo,
TableInfo,
ObjectInfo,
CompletionAssistantRequest,
CompletionAssistantResponse,
ObjectStatistics,
ObjectSource,
ObjectSourceKind,
@ -482,6 +484,10 @@ export async function listCompletionObjects(connectionId: string, database: stri
return get(`/api/schema/completion-objects?${qs({ connection_id: connectionId, database, schema })}`);
}
export async function completionAssistantSearch(request: CompletionAssistantRequest): Promise<CompletionAssistantResponse> {
return post("/api/schema/completion-assistant", request);
}
export async function getObjectSource(connectionId: string, database: string, schema: string, name: string, objectType: ObjectSourceKind): Promise<ObjectSource> {
return get(`/api/schema/object-source?${qs({ connection_id: connectionId, database, schema, table: name, object_type: objectType })}`);
}

View File

@ -1063,6 +1063,8 @@ export interface SqlCompletionReferencedTable {
export type SqlStatementKind = "select" | "insert" | "update" | "delete" | "create" | "alter" | "drop" | "unknown";
export type SqlCompletionContextKind = "table" | "schema" | "catalog" | "routine" | "column" | "alias_column" | "insert_target" | "update_target" | "exec" | "join" | "keyword";
export interface SqlCompletionContext {
prefix: string;
qualifier?: string;
@ -1090,6 +1092,7 @@ export interface SqlCompletionContext {
updateTarget?: { table: string; schema?: string };
deleteTarget?: { table: string; schema?: string };
oracleTableFunctionContext?: boolean;
contextKind: SqlCompletionContextKind;
}
export interface SqlFunctionSignatureHelp {
@ -1553,6 +1556,19 @@ export function getSqlCompletionContext(sql: string, cursor: number): SqlComplet
const statementKind = detectStatementKind(beforeCursor || fullStatement);
const preferredKeywords = preferredKeywordsForCompletion(updateInfo, deleteInfo);
const contextKind = detectCompletionContextKind({
qualifier,
exclusiveTableSuggestions,
exclusiveColumnSuggestions,
insertInfo,
updateInfo,
inCallRoutineContext,
oracleTableFunctionContext,
afterTableTrigger,
lastWord,
suggestColumns,
suggestRoutines,
});
return {
prefix,
@ -1581,9 +1597,33 @@ export function getSqlCompletionContext(sql: string, cursor: number): SqlComplet
updateTarget: updateInfo?.target,
deleteTarget: deleteInfo?.target,
oracleTableFunctionContext,
contextKind,
};
}
function detectCompletionContextKind(options: {
qualifier?: string;
exclusiveTableSuggestions: boolean;
exclusiveColumnSuggestions: boolean;
insertInfo: ReturnType<typeof detectInsertColumnListContext>;
updateInfo: ReturnType<typeof detectUpdateCompletionContext>;
inCallRoutineContext: boolean;
oracleTableFunctionContext: boolean;
afterTableTrigger: boolean;
lastWord: string;
suggestColumns: boolean;
suggestRoutines: boolean;
}): SqlCompletionContextKind {
if (options.insertInfo) return "column";
if (options.updateInfo?.inSetClause) return "column";
if (options.inCallRoutineContext) return "exec";
if (options.qualifier && options.exclusiveColumnSuggestions) return "alias_column";
if (options.oracleTableFunctionContext || options.suggestRoutines) return "routine";
if (options.exclusiveTableSuggestions || options.afterTableTrigger) return options.lastWord === "join" ? "join" : "table";
if (options.suggestColumns) return options.qualifier ? "alias_column" : "column";
return "keyword";
}
function parseTrailingIdentifierContext(input: string): { start: number; prefix: string; qualifier?: string; qualifierParts?: string[] } | null {
if (/\s$/.test(input)) return null;
let i = input.length - 1;

View File

@ -6,6 +6,8 @@ import type {
LinkedServerInfo,
TableInfo,
ObjectInfo,
CompletionAssistantRequest,
CompletionAssistantResponse,
ObjectStatistics,
ObjectSource,
ObjectSourceKind,
@ -533,6 +535,10 @@ export async function listCompletionObjects(connectionId: string, database: stri
return invoke("list_completion_objects", { connectionId, database, schema });
}
export async function completionAssistantSearch(request: CompletionAssistantRequest): Promise<CompletionAssistantResponse> {
return invoke("completion_assistant_search", { request });
}
export async function getObjectSource(connectionId: string, database: string, schema: string, name: string, objectType: ObjectSourceKind): Promise<ObjectSource> {
return invoke("get_object_source", { connectionId, database, schema, name, objectType });
}

View File

@ -0,0 +1,86 @@
import { createPinia, setActivePinia } from "pinia";
import { beforeEach, describe, expect, it, vi } from "vitest";
import type { ConnectionConfig } from "@/types/database";
function installLocalStorage() {
const data = new Map<string, string>();
vi.stubGlobal("localStorage", {
getItem: vi.fn((key: string) => data.get(key) ?? null),
setItem: vi.fn((key: string, value: string) => data.set(key, value)),
removeItem: vi.fn((key: string) => data.delete(key)),
});
}
function postgresConnection(): ConnectionConfig {
return {
id: "pg-1",
name: "Postgres",
db_type: "postgres",
host: "127.0.0.1",
port: 5432,
username: "postgres",
password: "",
database: "app",
read_only: false,
} as ConnectionConfig;
}
describe("connectionStore completion assistant", () => {
beforeEach(() => {
vi.resetModules();
vi.unstubAllGlobals();
installLocalStorage();
setActivePinia(createPinia());
});
it("deduplicates in-flight assistant table requests", async () => {
const completionAssistantSearch = vi.fn().mockResolvedValue({
candidates: [{ name: "accounts", kind: "table", schema: "public" }],
incomplete: false,
fallback_used: false,
});
vi.doMock("@/lib/tauriRuntime", () => ({ isTauriRuntime: () => false }));
vi.doMock("@/lib/api", () => ({
checkConnectionHealth: vi.fn().mockResolvedValue(undefined),
completionAssistantSearch,
listSchemas: vi.fn().mockResolvedValue(["public"]),
listTables: vi.fn().mockResolvedValue([]),
}));
const { useConnectionStore } = await import("@/stores/connectionStore");
const store = useConnectionStore();
store.connections = [postgresConnection()];
store.connectedIds.add("pg-1");
const [first, second] = await Promise.all([store.listCompletionTables("pg-1", "app", "acc", 20, "public"), store.listCompletionTables("pg-1", "app", "acc", 20, "public")]);
expect(completionAssistantSearch).toHaveBeenCalledTimes(1);
expect(first).toEqual(second);
expect(first[0]).toMatchObject({ name: "accounts", schema: "public", type: "table" });
});
it("returns fallback metadata when assistant table search fails", async () => {
const completionAssistantSearch = vi.fn().mockRejectedValue(new Error("assistant unavailable"));
const listTables = vi.fn().mockResolvedValue([{ name: "accounts", table_type: "BASE TABLE", comment: null }]);
vi.doMock("@/lib/tauriRuntime", () => ({ isTauriRuntime: () => false }));
vi.doMock("@/lib/api", () => ({
checkConnectionHealth: vi.fn().mockResolvedValue(undefined),
completionAssistantSearch,
listSchemas: vi.fn().mockResolvedValue(["public"]),
listTables,
}));
const { useConnectionStore } = await import("@/stores/connectionStore");
const store = useConnectionStore();
store.connections = [postgresConnection()];
store.connectedIds.add("pg-1");
const tables = await store.listCompletionTables("pg-1", "app", "acc", 20, "public");
expect(completionAssistantSearch).toHaveBeenCalledTimes(1);
expect(listTables).toHaveBeenCalledWith("pg-1", "app", "public", "acc", 20);
expect(tables).toEqual([{ name: "accounts", schema: "public", type: "table" }]);
});
});

View File

@ -1,7 +1,7 @@
import { defineStore } from "pinia";
import { uuid } from "@/lib/utils";
import { ref, computed, watch } from "vue";
import type { ColumnInfo, ConnectionConfig, ForeignKeyInfo, ObjectInfo, SidebarLayout, TableInfo, TreeNode } from "@/types/database";
import type { ColumnInfo, CompletionAssistantCandidate, CompletionAssistantObjectKind, CompletionAssistantRequest, ConnectionConfig, ForeignKeyInfo, ObjectInfo, SidebarLayout, TableInfo, TreeNode } from "@/types/database";
import { applyPinnedTreeNodeState, updatePinnedTreeNodeInPlace } from "@/lib/pinnedItems";
import {
reconcileLayout,
@ -2143,6 +2143,87 @@ export const useConnectionStore = defineStore("connection", () => {
return promise;
}
function completionAssistantRequestKey(request: CompletionAssistantRequest): string {
return JSON.stringify({
connection_id: request.connection_id,
database: request.database,
schema: request.schema ?? "",
object_kinds: [...(request.object_kinds ?? [])].sort(),
mask: request.mask ?? "",
case_sensitive: !!request.case_sensitive,
global_search: !!request.global_search,
max_results: request.max_results ?? null,
search_in_comments: !!request.search_in_comments,
search_in_definitions: !!request.search_in_definitions,
parent_schema: request.parent_schema ?? "",
parent_name: request.parent_name ?? "",
match_mode: request.match_mode ?? "prefix",
});
}
async function completionAssistantSearch(request: CompletionAssistantRequest) {
return withCompletionInFlight(`assistant:${completionAssistantRequestKey(request)}`, async () => {
await ensureConnected(request.connection_id);
return api.completionAssistantSearch(request);
});
}
function completionAssistantTables(candidates: CompletionAssistantCandidate[]): SqlCompletionTable[] {
return candidates
.filter((candidate) => candidate.kind === "table" || candidate.kind === "view")
.map((candidate) => ({
name: candidate.name,
schema: candidate.schema ?? undefined,
type: candidate.kind === "view" ? ("view" as const) : ("table" as const),
}));
}
function completionAssistantColumns(candidates: CompletionAssistantCandidate[], table: string, schema?: string): SqlCompletionColumn[] {
return candidates
.filter((candidate) => candidate.kind === "column")
.map((candidate) => ({
name: candidate.name,
table: candidate.parent_name ?? table,
schema: candidate.parent_schema ?? candidate.schema ?? schema,
dataType: candidate.data_type ?? undefined,
comment: candidate.comment ?? null,
}));
}
async function listCompletionAssistantTables(connectionId: string, database: string, filter: string, limit?: number, schema?: string): Promise<SqlCompletionTable[]> {
const objectKinds: CompletionAssistantObjectKind[] = ["table", "view"];
const response = await completionAssistantSearch({
connection_id: connectionId,
database,
schema: schema ?? null,
object_kinds: objectKinds,
mask: filter.trim(),
max_results: limit ?? 200,
parent_schema: schema ?? null,
match_mode: "prefix",
});
const tables = completionAssistantTables(response.candidates);
indexCompletionTables(connectionId, database, schema, tables);
return tables;
}
async function listCompletionAssistantColumns(connectionId: string, database: string, table: string, schema?: string): Promise<SqlCompletionColumn[]> {
const response = await completionAssistantSearch({
connection_id: connectionId,
database,
schema: schema ?? null,
object_kinds: ["column"],
mask: "",
max_results: 500,
parent_schema: schema ?? null,
parent_name: table,
match_mode: "prefix",
});
const columns = completionAssistantColumns(response.candidates, table, schema);
if (columns.length > 0) indexCompletionColumns(connectionId, database, table, schema, columns);
return columns;
}
function completionNameSegments(name: string): string[] {
return name
.replace(/([a-z0-9])([A-Z])/g, "$1 $2")
@ -2432,53 +2513,36 @@ export const useConnectionStore = defineStore("connection", () => {
await ensureConnected(connectionId);
if (isSchemaAwareDatabase(connectionId)) {
const schemas = schema ? [schema] : await listCompletionSchemas(connectionId, database);
if (normalizedFilter || limit) {
const batchSize = 5;
const results: SqlCompletionTable[] = [];
const maxResults = limit ?? Infinity;
for (let i = 0; i < schemas.length && results.length < maxResults; i += batchSize) {
const batch = schemas.slice(i, i + batchSize);
const batchResults = await Promise.all(
batch.map(async (s) => {
try {
const tables = await api.listTables(connectionId, database, s, normalizedFilter, limit);
return tables.map((table) => ({
name: table.name,
schema: s,
type: table.table_type === "VIEW" || table.table_type === "MATERIALIZED_VIEW" ? ("view" as const) : ("table" as const),
})) as SqlCompletionTable[];
} catch {
return [] as SqlCompletionTable[];
}
}),
);
for (const group of batchResults) {
results.push(...group);
indexCompletionTables(connectionId, database, undefined, group);
let results: SqlCompletionTable[] = [];
try {
results = await listCompletionAssistantTables(connectionId, database, normalizedFilter, limit, schema);
} catch {
if (schema) {
const tables = await api.listTables(connectionId, database, schema, normalizedFilter, limit);
results = tables.map((table) => ({
name: table.name,
schema,
type: table.table_type === "VIEW" || table.table_type === "MATERIALIZED_VIEW" ? ("view" as const) : ("table" as const),
}));
} else {
results = lookupLocalCompletionTables(connectionId, database, normalizedFilter, limit);
}
}
if (results.length === 0 && relaxedFilter) {
for (let i = 0; i < schemas.length && results.length < maxResults; i += batchSize) {
const batch = schemas.slice(i, i + batchSize);
const batchResults = await Promise.all(
batch.map(async (s) => {
try {
const tables = await api.listTables(connectionId, database, s, relaxedFilter, expandedCompletionLimit(limit));
return tables.map((table) => ({
name: table.name,
schema: s,
type: table.table_type === "VIEW" || table.table_type === "MATERIALIZED_VIEW" ? ("view" as const) : ("table" as const),
})) as SqlCompletionTable[];
} catch {
return [] as SqlCompletionTable[];
}
}),
);
for (const group of batchResults) {
results.push(...group);
indexCompletionTables(connectionId, database, undefined, group);
if (schema) {
try {
const tables = await api.listTables(connectionId, database, schema, relaxedFilter, expandedCompletionLimit(limit));
results = tables.map((table) => ({
name: table.name,
schema,
type: table.table_type === "VIEW" || table.table_type === "MATERIALIZED_VIEW" ? ("view" as const) : ("table" as const),
}));
} catch {
results = [];
}
} else {
results = lookupLocalCompletionTables(connectionId, database, relaxedFilter, expandedCompletionLimit(limit));
}
}
const limitedTables = limit ? dedupeCompletionTables(results).slice(0, limit) : results;
@ -2488,21 +2552,16 @@ export const useConnectionStore = defineStore("connection", () => {
return completionTablesCache.value[cacheKey];
}
const tableGroups = await Promise.all(
schemas.map(async (schema) => {
try {
const tables = await api.listTables(connectionId, database, schema);
return tables.map((table) => ({
name: table.name,
schema,
type: table.table_type === "VIEW" || table.table_type === "MATERIALIZED_VIEW" ? ("view" as const) : ("table" as const),
}));
} catch {
return [];
}
}),
);
completionTablesCache.value[cacheKey] = tableGroups.flat();
if (schema) {
const tables = await api.listTables(connectionId, database, schema);
completionTablesCache.value[cacheKey] = tables.map((table) => ({
name: table.name,
schema,
type: table.table_type === "VIEW" || table.table_type === "MATERIALIZED_VIEW" ? ("view" as const) : ("table" as const),
}));
} else {
completionTablesCache.value[cacheKey] = lookupLocalCompletionTables(connectionId, database, normalizedFilter, limit);
}
indexCompletionTables(connectionId, database, undefined, completionTablesCache.value[cacheKey]);
evictOldestCacheEntries(completionTablesCache.value, COMPLETION_CACHE_MAX);
return completionTablesCache.value[cacheKey];
@ -2633,6 +2692,27 @@ export const useConnectionStore = defineStore("connection", () => {
if (!completionColumnsCache.value[cacheKey]) {
await withCompletionInFlight(`${cacheKey}:columns`, async () => {
await ensureConnected(connectionId);
try {
const assistantColumns = await listCompletionAssistantColumns(connectionId, database, table, schema);
if (assistantColumns.length > 0) {
completionColumnsCache.value[cacheKey] = assistantColumns.map((column) => ({
name: column.name,
data_type: column.dataType ?? "",
is_nullable: column.isNullable ?? true,
column_default: null,
is_primary_key: false,
extra: null,
comment: column.comment ?? null,
numeric_precision: null,
numeric_scale: null,
character_maximum_length: null,
}));
evictOldestCacheEntries(completionColumnsCache.value, COMPLETION_CACHE_MAX);
return;
}
} catch {
// Fall back to the existing metadata path below.
}
const querySchema = metadataQuerySchema(connectionId, database, schema);
completionColumnsCache.value[cacheKey] = await api.getColumns(connectionId, database, querySchema, table);
evictOldestCacheEntries(completionColumnsCache.value, COMPLETION_CACHE_MAX);

View File

@ -65,6 +65,45 @@ export interface SqlSnippet {
body: string;
}
export type CompletionAssistantObjectKind = "database" | "schema" | "table" | "view" | "routine" | "procedure" | "function" | "column";
export type CompletionAssistantCandidateKind = "database" | "schema" | "table" | "view" | "procedure" | "function" | "column" | "object";
export type CompletionAssistantMatchMode = "prefix" | "contains";
export interface CompletionAssistantRequest {
connection_id: string;
database: string;
schema?: string | null;
object_kinds?: CompletionAssistantObjectKind[];
mask?: string;
case_sensitive?: boolean;
global_search?: boolean;
max_results?: number | null;
search_in_comments?: boolean;
search_in_definitions?: boolean;
parent_schema?: string | null;
parent_name?: string | null;
match_mode?: CompletionAssistantMatchMode | null;
}
export interface CompletionAssistantCandidate {
name: string;
kind: CompletionAssistantCandidateKind;
database?: string | null;
schema?: string | null;
parent_schema?: string | null;
parent_name?: string | null;
comment?: string | null;
data_type?: string | null;
}
export interface CompletionAssistantResponse {
candidates: CompletionAssistantCandidate[];
incomplete: boolean;
fallback_used: boolean;
}
export interface ConnectionConfig {
id: string;
name: string;

View File

@ -43,6 +43,7 @@
"list_schemas",
"list_tables",
"list_objects",
"completion_assistant_search_v1",
"get_object_source",
"get_table_ddl",
"get_columns",

View File

@ -120,6 +120,7 @@ pub enum AgentMethod {
ListSchemas,
ListTables,
ListObjects,
CompletionAssistantSearchV1,
GetObjectSource,
GetColumns,
ListIndexes,
@ -137,7 +138,7 @@ pub enum AgentMethod {
}
impl AgentMethod {
pub const ALL: [Self; 22] = [
pub const ALL: [Self; 23] = [
Self::Handshake,
Self::Connect,
Self::TestConnection,
@ -146,6 +147,7 @@ impl AgentMethod {
Self::ListSchemas,
Self::ListTables,
Self::ListObjects,
Self::CompletionAssistantSearchV1,
Self::GetObjectSource,
Self::GetTableDdl,
Self::GetColumns,
@ -172,6 +174,7 @@ impl AgentMethod {
Self::ListSchemas => "list_schemas",
Self::ListTables => "list_tables",
Self::ListObjects => "list_objects",
Self::CompletionAssistantSearchV1 => "completion_assistant_search_v1",
Self::GetObjectSource => "get_object_source",
Self::GetTableDdl => "get_table_ddl",
Self::GetColumns => "get_columns",
@ -569,6 +572,19 @@ impl AgentDriverClient {
.await
}
pub async fn completion_assistant_search<T: DeserializeOwned + Send + 'static>(
&mut self,
request: &crate::types::CompletionAssistantRequest,
timeout_duration: Option<Duration>,
) -> Result<T, String> {
self.call_method_with_timeout(
AgentMethod::CompletionAssistantSearchV1,
serde_json::to_value(request).map_err(|e| e.to_string())?,
timeout_duration,
)
.await
}
pub async fn get_object_source<T: DeserializeOwned + Send + 'static, K: Serialize>(
&mut self,
database: &str,
@ -1245,6 +1261,7 @@ mod tests {
assert_eq!(AgentMethod::ListSchemas.as_str(), "list_schemas");
assert_eq!(AgentMethod::ListTables.as_str(), "list_tables");
assert_eq!(AgentMethod::ListObjects.as_str(), "list_objects");
assert_eq!(AgentMethod::CompletionAssistantSearchV1.as_str(), "completion_assistant_search_v1");
assert_eq!(AgentMethod::GetObjectSource.as_str(), "get_object_source");
assert_eq!(AgentMethod::GetColumns.as_str(), "get_columns");
assert_eq!(AgentMethod::ListIndexes.as_str(), "list_indexes");

View File

@ -13,8 +13,9 @@ use std::time::Instant;
use crate::models::connection::DatabaseType;
use crate::sql::starts_with_executable_sql_keyword;
use crate::types::{
ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, ObjectInfo, ObjectStatistics, QueryResult, TableInfo,
TriggerInfo,
ColumnInfo, CompletionAssistantCandidate, CompletionAssistantCandidateKind, CompletionAssistantMatchMode,
CompletionAssistantObjectKind, CompletionAssistantRequest, CompletionAssistantResponse, DatabaseInfo,
ForeignKeyInfo, IndexInfo, ObjectInfo, ObjectStatistics, QueryResult, TableInfo, TriggerInfo,
};
use super::file_validator::validate_file_path;
@ -1025,6 +1026,228 @@ pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result<Vec<TableIn
Ok(tables)
}
pub async fn completion_assistant_search(
pool: &MySqlPool,
request: &CompletionAssistantRequest,
) -> Result<CompletionAssistantResponse, String> {
let database = request.schema.as_deref().filter(|schema| !schema.trim().is_empty()).unwrap_or(&request.database);
let limit = request.max_results.unwrap_or(100).clamp(1, 1000);
let kinds = if request.object_kinds.is_empty() {
vec![CompletionAssistantObjectKind::Table, CompletionAssistantObjectKind::View]
} else {
request.object_kinds.clone()
};
let pattern = mysql_completion_like_pattern(&request.mask, request.match_mode.as_ref());
let mut conn = pool.get_conn().await.map_err(|e| e.to_string())?;
let mut candidates = Vec::new();
if kinds
.iter()
.any(|kind| matches!(kind, CompletionAssistantObjectKind::Database | CompletionAssistantObjectKind::Schema))
{
let sql = mysql_completion_schemas_sql(&pattern, limit.saturating_sub(candidates.len()));
let result = conn.query_iter(&sql).await.map_err(|e| e.to_string())?;
let rows: Vec<mysql_async::Row> = result.collect_and_drop().await.map_err(|e| e.to_string())?;
for row in rows {
let schema_name = get_str_by_name(&row, "schema_name");
candidates.push(CompletionAssistantCandidate {
name: schema_name.clone(),
kind: CompletionAssistantCandidateKind::Schema,
database: Some(schema_name.clone()),
schema: Some(schema_name),
parent_schema: None,
parent_name: None,
comment: None,
data_type: None,
});
}
}
if candidates.len() < limit && kinds.iter().any(CompletionAssistantObjectKind::is_table_like) {
let sql = mysql_completion_tables_sql(database, &pattern, &kinds, limit.saturating_sub(candidates.len()));
let result = conn.query_iter(&sql).await.map_err(|e| e.to_string())?;
let rows: Vec<mysql_async::Row> = result.collect_and_drop().await.map_err(|e| e.to_string())?;
for row in rows {
let table_type = get_str_by_name(&row, "table_type");
candidates.push(CompletionAssistantCandidate {
name: get_str_by_name(&row, "object_name"),
kind: if table_type.eq_ignore_ascii_case("VIEW") {
CompletionAssistantCandidateKind::View
} else {
CompletionAssistantCandidateKind::Table
},
database: Some(database.to_string()),
schema: Some(database.to_string()),
parent_schema: None,
parent_name: None,
comment: get_opt_str(&row, "object_comment")
.map(|s| fix_potential_double_encoding(&s))
.filter(|s| !s.is_empty()),
data_type: None,
});
}
}
if candidates.len() < limit && kinds.iter().any(CompletionAssistantObjectKind::is_routine_like) {
let sql = mysql_completion_routines_sql(database, &pattern, &kinds, limit.saturating_sub(candidates.len()));
let result = conn.query_iter(&sql).await.map_err(|e| e.to_string())?;
let rows: Vec<mysql_async::Row> = result.collect_and_drop().await.map_err(|e| e.to_string())?;
for row in rows {
let routine_type = get_str_by_name(&row, "routine_type");
candidates.push(CompletionAssistantCandidate {
name: get_str_by_name(&row, "object_name"),
kind: if routine_type.eq_ignore_ascii_case("PROCEDURE") {
CompletionAssistantCandidateKind::Procedure
} else {
CompletionAssistantCandidateKind::Function
},
database: Some(database.to_string()),
schema: Some(database.to_string()),
parent_schema: None,
parent_name: None,
comment: get_opt_str(&row, "object_comment")
.map(|s| fix_potential_double_encoding(&s))
.filter(|s| !s.is_empty()),
data_type: get_opt_str(&row, "data_type"),
});
}
}
if candidates.len() < limit && kinds.iter().any(|kind| matches!(kind, CompletionAssistantObjectKind::Column)) {
if let Some(table) = request.parent_name.as_deref().filter(|table| !table.trim().is_empty()) {
let sql = mysql_completion_columns_sql(database, table, &pattern, limit.saturating_sub(candidates.len()));
let result = conn.query_iter(&sql).await.map_err(|e| e.to_string())?;
let rows: Vec<mysql_async::Row> = result.collect_and_drop().await.map_err(|e| e.to_string())?;
for row in rows {
candidates.push(CompletionAssistantCandidate {
name: get_str_by_name(&row, "object_name"),
kind: CompletionAssistantCandidateKind::Column,
database: Some(database.to_string()),
schema: Some(database.to_string()),
parent_schema: Some(database.to_string()),
parent_name: Some(table.to_string()),
comment: get_opt_str(&row, "object_comment")
.map(|s| fix_potential_double_encoding(&s))
.filter(|s| !s.is_empty()),
data_type: Some(get_str_by_name(&row, "data_type")),
});
}
}
}
Ok(CompletionAssistantResponse { incomplete: candidates.len() >= limit, candidates, fallback_used: false })
}
fn mysql_completion_schemas_sql(pattern: &str, limit: usize) -> String {
format!(
"SELECT SCHEMA_NAME AS schema_name \
FROM information_schema.SCHEMATA \
WHERE SCHEMA_NAME LIKE {} ESCAPE '\\\\' \
ORDER BY SCHEMA_NAME LIMIT {}",
quote_value(pattern),
limit,
)
}
fn mysql_completion_tables_sql(
database: &str,
pattern: &str,
kinds: &[CompletionAssistantObjectKind],
limit: usize,
) -> String {
let table_types = mysql_completion_table_types(kinds);
format!(
"SELECT TABLE_NAME AS object_name, TABLE_TYPE AS table_type, TABLE_COMMENT AS object_comment \
FROM information_schema.TABLES \
WHERE TABLE_SCHEMA = {db} AND TABLE_NAME LIKE {pattern} ESCAPE '\\\\' AND TABLE_TYPE IN ({table_types}) \
ORDER BY TABLE_NAME LIMIT {limit}",
db = quote_value(database),
pattern = quote_value(pattern),
table_types = table_types,
limit = limit,
)
}
fn mysql_completion_routines_sql(
database: &str,
pattern: &str,
kinds: &[CompletionAssistantObjectKind],
limit: usize,
) -> String {
let routine_types = mysql_completion_routine_types(kinds);
format!(
"SELECT ROUTINE_NAME AS object_name, ROUTINE_TYPE AS routine_type, ROUTINE_COMMENT AS object_comment, DATA_TYPE AS data_type \
FROM information_schema.ROUTINES \
WHERE ROUTINE_SCHEMA = {db} AND ROUTINE_NAME LIKE {pattern} ESCAPE '\\\\' AND ROUTINE_TYPE IN ({routine_types}) \
ORDER BY ROUTINE_NAME LIMIT {limit}",
db = quote_value(database),
pattern = quote_value(pattern),
routine_types = routine_types,
limit = limit,
)
}
fn mysql_completion_columns_sql(database: &str, table: &str, pattern: &str, limit: usize) -> String {
format!(
"SELECT COLUMN_NAME AS object_name, COLUMN_TYPE AS data_type, COLUMN_COMMENT AS object_comment \
FROM information_schema.COLUMNS \
WHERE TABLE_SCHEMA = {db} AND TABLE_NAME = {table} AND COLUMN_NAME LIKE {pattern} ESCAPE '\\\\' \
ORDER BY ORDINAL_POSITION LIMIT {limit}",
db = quote_value(database),
table = quote_value(table),
pattern = quote_value(pattern),
limit = limit,
)
}
fn mysql_completion_table_types(kinds: &[CompletionAssistantObjectKind]) -> String {
let mut types = Vec::new();
if kinds.iter().any(|kind| matches!(kind, CompletionAssistantObjectKind::Table)) {
types.push("'BASE TABLE'");
types.push("'SYSTEM VERSIONED'");
}
if kinds.iter().any(|kind| matches!(kind, CompletionAssistantObjectKind::View)) {
types.push("'VIEW'");
}
if types.is_empty() {
"'BASE TABLE','VIEW'".to_string()
} else {
types.join(",")
}
}
fn mysql_completion_routine_types(kinds: &[CompletionAssistantObjectKind]) -> String {
let mut types = Vec::new();
if kinds
.iter()
.any(|kind| matches!(kind, CompletionAssistantObjectKind::Procedure | CompletionAssistantObjectKind::Routine))
{
types.push("'PROCEDURE'");
}
if kinds
.iter()
.any(|kind| matches!(kind, CompletionAssistantObjectKind::Function | CompletionAssistantObjectKind::Routine))
{
types.push("'FUNCTION'");
}
if types.is_empty() {
"'PROCEDURE','FUNCTION'".to_string()
} else {
types.join(",")
}
}
fn mysql_completion_like_pattern(value: &str, mode: Option<&CompletionAssistantMatchMode>) -> String {
if value.trim().is_empty() || value == "%" {
return "%".to_string();
}
let escaped = value.trim().replace('\\', "\\\\").replace('%', "\\%").replace('_', "\\_");
match mode.unwrap_or(&CompletionAssistantMatchMode::Prefix) {
CompletionAssistantMatchMode::Prefix => format!("{escaped}%"),
CompletionAssistantMatchMode::Contains => format!("%{escaped}%"),
}
}
fn table_comment_sql(database: &str, table: &str) -> String {
format!(
"SELECT TABLE_COMMENT \
@ -2153,6 +2376,37 @@ mod tests {
assert!(sql.contains("TRIGGER_SCHEMA = 'app'"));
}
#[test]
fn mysql_completion_like_pattern_uses_prefix_by_default() {
assert_eq!(mysql_completion_like_pattern("Temp", Some(&CompletionAssistantMatchMode::Prefix)), "Temp%");
assert_eq!(mysql_completion_like_pattern("Temp", Some(&CompletionAssistantMatchMode::Contains)), "%Temp%");
assert_eq!(
mysql_completion_like_pattern("order_100%", Some(&CompletionAssistantMatchMode::Prefix)),
"order\\_100\\%%"
);
}
#[test]
fn mysql_completion_sql_filters_before_limit() {
let table_sql = mysql_completion_tables_sql(
"app",
"Temp%",
&[CompletionAssistantObjectKind::Table, CompletionAssistantObjectKind::View],
100,
);
let routine_sql =
mysql_completion_routines_sql("app", "%audit%", &[CompletionAssistantObjectKind::Routine], 50);
let column_sql = mysql_completion_columns_sql("app", "users", "id%", 25);
assert!(table_sql.contains("TABLE_NAME LIKE 'Temp%' ESCAPE '\\\\'"));
assert!(table_sql.contains("TABLE_TYPE IN ('BASE TABLE','SYSTEM VERSIONED','VIEW')"));
assert!(table_sql.contains("ORDER BY TABLE_NAME LIMIT 100"));
assert!(routine_sql.contains("ROUTINE_NAME LIKE '%audit%' ESCAPE '\\\\'"));
assert!(routine_sql.contains("ROUTINE_TYPE IN ('PROCEDURE','FUNCTION')"));
assert!(column_sql.contains("COLUMN_NAME LIKE 'id%' ESCAPE '\\\\'"));
assert!(column_sql.contains("ORDER BY ORDINAL_POSITION LIMIT 25"));
}
#[test]
fn mysql_columns_sql_uses_column_key_for_primary_keys_without_join() {
let sql = columns_sql("app", "users");

View File

@ -22,8 +22,10 @@ use tokio_util::sync::CancellationToken;
use super::file_validator::validate_file_path;
use crate::sql::starts_with_executable_sql_keyword;
use crate::types::{
ColumnInfo, DatabaseInfo, ForeignKeyInfo, FunctionInfo, IndexInfo, ObjectInfo, ObjectStatistics, OwnerInfo,
QueryResult, RuleInfo, SequenceInfo, TableInfo, TriggerInfo,
ColumnInfo, CompletionAssistantCandidate, CompletionAssistantCandidateKind, CompletionAssistantMatchMode,
CompletionAssistantObjectKind, CompletionAssistantRequest, CompletionAssistantResponse, DatabaseInfo,
ForeignKeyInfo, FunctionInfo, IndexInfo, ObjectInfo, ObjectStatistics, OwnerInfo, QueryResult, RuleInfo,
SequenceInfo, TableInfo, TriggerInfo,
};
fn pg_temporal_to_json_value(row: &Row, idx: usize) -> Option<serde_json::Value> {
@ -964,6 +966,199 @@ pub async fn list_tables_filtered(
.collect())
}
pub async fn completion_assistant_search(
pool: &Pool,
request: &CompletionAssistantRequest,
) -> Result<CompletionAssistantResponse, String> {
let schema = request.schema.as_deref().or(request.parent_schema.as_deref()).unwrap_or("public");
let limit = request.max_results.unwrap_or(100).clamp(1, 1000);
let kinds = if request.object_kinds.is_empty() {
vec![CompletionAssistantObjectKind::Table, CompletionAssistantObjectKind::View]
} else {
request.object_kinds.clone()
};
let pattern = postgres_completion_like_pattern(&request.mask, request.match_mode.as_ref());
let client = pool.get().await.map_err(|e| e.to_string())?;
let mut candidates = Vec::new();
if kinds.iter().any(|kind| matches!(kind, CompletionAssistantObjectKind::Schema)) {
let stmt = client
.prepare_cached(
"SELECT nspname FROM pg_catalog.pg_namespace \
WHERE nspname NOT LIKE 'pg_%' AND nspname <> 'information_schema' \
AND ($1 = '%%' OR nspname ILIKE $1 ESCAPE '~') \
ORDER BY nspname LIMIT $2",
)
.await
.map_err(|e| e.to_string())?;
for row in client.query(&stmt, &[&pattern, &(limit as i64)]).await.map_err(|e| e.to_string())? {
let schema_name: String = row.get(0);
candidates.push(CompletionAssistantCandidate {
name: schema_name.clone(),
kind: CompletionAssistantCandidateKind::Schema,
database: Some(request.database.clone()),
schema: Some(schema_name),
parent_schema: None,
parent_name: None,
comment: None,
data_type: None,
});
}
}
if candidates.len() < limit && kinds.iter().any(CompletionAssistantObjectKind::is_table_like) {
let relkinds = postgres_completion_relkinds(&kinds);
let stmt = client.prepare_cached(postgres_completion_tables_sql()).await.map_err(|e| e.to_string())?;
let rows = client
.query(&stmt, &[&schema, &pattern, &relkinds, &((limit - candidates.len()) as i64)])
.await
.map_err(|e| e.to_string())?;
for row in rows {
let table_type: String = row.get(2);
candidates.push(CompletionAssistantCandidate {
name: row.get(0),
kind: if table_type == "VIEW" {
CompletionAssistantCandidateKind::View
} else {
CompletionAssistantCandidateKind::Table
},
database: Some(request.database.clone()),
schema: Some(row.get(1)),
parent_schema: row.try_get::<_, Option<String>>(4).ok().flatten(),
parent_name: row.try_get::<_, Option<String>>(5).ok().flatten(),
comment: row.try_get::<_, Option<String>>(3).ok().flatten(),
data_type: None,
});
}
}
if candidates.len() < limit && kinds.iter().any(CompletionAssistantObjectKind::is_routine_like) {
let prokinds = postgres_completion_prokinds(&kinds);
let stmt = client.prepare_cached(postgres_completion_routines_sql()).await.map_err(|e| e.to_string())?;
let rows = client
.query(&stmt, &[&schema, &pattern, &prokinds, &((limit - candidates.len()) as i64)])
.await
.map_err(|e| e.to_string())?;
for row in rows {
let routine_type: String = row.get(2);
candidates.push(CompletionAssistantCandidate {
name: row.get(0),
kind: if routine_type == "PROCEDURE" {
CompletionAssistantCandidateKind::Procedure
} else {
CompletionAssistantCandidateKind::Function
},
database: Some(request.database.clone()),
schema: Some(row.get(1)),
parent_schema: None,
parent_name: None,
comment: row.try_get::<_, Option<String>>(3).ok().flatten(),
data_type: row.try_get::<_, Option<String>>(4).ok().flatten(),
});
}
}
if candidates.len() < limit && kinds.iter().any(|kind| matches!(kind, CompletionAssistantObjectKind::Column)) {
let table = request.parent_name.as_deref().unwrap_or("");
if !table.is_empty() {
let stmt = client.prepare_cached(postgres_completion_columns_sql()).await.map_err(|e| e.to_string())?;
let rows = client
.query(&stmt, &[&schema, &table, &pattern, &((limit - candidates.len()) as i64)])
.await
.map_err(|e| e.to_string())?;
for row in rows {
candidates.push(CompletionAssistantCandidate {
name: row.get(0),
kind: CompletionAssistantCandidateKind::Column,
database: Some(request.database.clone()),
schema: Some(schema.to_string()),
parent_schema: Some(schema.to_string()),
parent_name: Some(table.to_string()),
comment: row.try_get::<_, Option<String>>(2).ok().flatten(),
data_type: Some(row.get(1)),
});
}
}
}
Ok(CompletionAssistantResponse { incomplete: candidates.len() >= limit, candidates, fallback_used: false })
}
fn postgres_completion_tables_sql() -> &'static str {
"SELECT c.relname, n.nspname, \
CASE c.relkind WHEN 'v' THEN 'VIEW' WHEN 'm' THEN 'VIEW' ELSE 'TABLE' END AS table_type, \
obj_description(c.oid) AS table_comment, \
CASE WHEN pc.relkind = 'p' THEN pn.nspname ELSE NULL END AS parent_schema, \
CASE WHEN pc.relkind = 'p' THEN pc.relname ELSE NULL END AS parent_name \
FROM pg_catalog.pg_class c \
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace \
LEFT JOIN pg_catalog.pg_inherits i ON i.inhrelid = c.oid \
LEFT JOIN pg_catalog.pg_class pc ON pc.oid = i.inhparent \
LEFT JOIN pg_catalog.pg_namespace pn ON pn.oid = pc.relnamespace \
WHERE n.nspname = $1 AND c.relkind = ANY($3) \
AND ($2 = '%%' OR c.relname ILIKE $2 ESCAPE '~') \
ORDER BY c.relname LIMIT $4"
}
fn postgres_completion_routines_sql() -> &'static str {
"SELECT p.proname, n.nspname, CASE p.prokind WHEN 'p' THEN 'PROCEDURE' ELSE 'FUNCTION' END, \
obj_description(p.oid) AS routine_comment, COALESCE(pg_get_function_result(p.oid), '') AS data_type \
FROM pg_catalog.pg_proc p \
JOIN pg_catalog.pg_namespace n ON n.oid = p.pronamespace \
WHERE n.nspname = $1 AND p.prokind = ANY($3) \
AND ($2 = '%%' OR p.proname ILIKE $2 ESCAPE '~') \
ORDER BY p.proname LIMIT $4"
}
fn postgres_completion_columns_sql() -> &'static str {
"SELECT a.attname, pg_catalog.format_type(a.atttypid, a.atttypmod), col_description(c.oid, a.attnum) \
FROM pg_catalog.pg_attribute a \
JOIN pg_catalog.pg_class c ON c.oid = a.attrelid \
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace \
WHERE n.nspname = $1 AND c.relname = $2 AND a.attnum > 0 AND NOT a.attisdropped \
AND ($3 = '%%' OR a.attname ILIKE $3 ESCAPE '~') \
ORDER BY a.attnum LIMIT $4"
}
fn postgres_completion_relkinds(kinds: &[CompletionAssistantObjectKind]) -> Vec<String> {
let mut relkinds = Vec::new();
if kinds.iter().any(|kind| matches!(kind, CompletionAssistantObjectKind::Table)) {
relkinds.extend(["r", "p", "f"].into_iter().map(str::to_string));
}
if kinds.iter().any(|kind| matches!(kind, CompletionAssistantObjectKind::View)) {
relkinds.extend(["v", "m"].into_iter().map(str::to_string));
}
relkinds
}
fn postgres_completion_prokinds(kinds: &[CompletionAssistantObjectKind]) -> Vec<String> {
let mut prokinds = Vec::new();
if kinds
.iter()
.any(|kind| matches!(kind, CompletionAssistantObjectKind::Procedure | CompletionAssistantObjectKind::Routine))
{
prokinds.push("p".to_string());
}
if kinds
.iter()
.any(|kind| matches!(kind, CompletionAssistantObjectKind::Function | CompletionAssistantObjectKind::Routine))
{
prokinds.push("f".to_string());
}
prokinds
}
fn postgres_completion_like_pattern(value: &str, mode: Option<&CompletionAssistantMatchMode>) -> String {
if value.trim().is_empty() || value == "%" {
return "%%".to_string();
}
let escaped = value.trim().replace('~', "~~").replace('%', "~%").replace('_', "~_");
match mode.unwrap_or(&CompletionAssistantMatchMode::Prefix) {
CompletionAssistantMatchMode::Prefix => format!("{escaped}%"),
CompletionAssistantMatchMode::Contains => format!("%{escaped}%"),
}
}
pub async fn get_table_comment(pool: &Pool, schema: &str, table: &str) -> Result<Option<String>, String> {
let schema = if schema.is_empty() { "public" } else { schema };
let client = pool.get().await.map_err(|e| e.to_string())?;
@ -2526,4 +2721,22 @@ mod tests {
assert!(sql.contains("ILIKE $2 ESCAPE '~'"));
}
#[test]
fn postgres_completion_like_pattern_uses_prefix_by_default() {
assert_eq!(postgres_completion_like_pattern("Temp", Some(&CompletionAssistantMatchMode::Prefix)), "Temp%");
assert_eq!(postgres_completion_like_pattern("Temp", Some(&CompletionAssistantMatchMode::Contains)), "%Temp%");
assert_eq!(
postgres_completion_like_pattern("order_100%", Some(&CompletionAssistantMatchMode::Prefix)),
"order~_100~%%"
);
}
#[test]
fn postgres_completion_sql_filters_before_limit() {
assert!(postgres_completion_tables_sql().contains("c.relname ILIKE $2 ESCAPE '~'"));
assert!(postgres_completion_tables_sql().contains("ORDER BY c.relname LIMIT $4"));
assert!(postgres_completion_routines_sql().contains("p.proname ILIKE $2 ESCAPE '~'"));
assert!(postgres_completion_columns_sql().contains("a.attname ILIKE $3 ESCAPE '~'"));
}
}

View File

@ -9,7 +9,11 @@ use std::time::Instant;
use super::file_validator::validate_file_path;
use crate::sql::starts_with_executable_sql_keyword;
use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo};
use crate::types::{
ColumnInfo, CompletionAssistantCandidate, CompletionAssistantCandidateKind, CompletionAssistantMatchMode,
CompletionAssistantObjectKind, CompletionAssistantRequest, CompletionAssistantResponse, DatabaseInfo,
ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo,
};
const SQLITE_DATABASE_HEADER: &[u8; 16] = b"SQLite format 3\0";
@ -567,6 +571,68 @@ mod tests {
let id = cols.iter().find(|c| c.name == "id").expect("id col");
assert!(id.extra.is_none());
}
#[tokio::test]
async fn completion_assistant_searches_sqlite_tables_and_columns_with_limit() {
let pool = connect_path(":memory:").await.expect("connect in-memory SQLite");
execute_query(
&pool,
"CREATE TABLE account(id INTEGER PRIMARY KEY, display_name TEXT); CREATE VIEW account_view AS SELECT id FROM account; CREATE TABLE audit_log(id INTEGER);",
)
.await
.expect("setup schema");
let tables = completion_assistant_search(
&pool,
&CompletionAssistantRequest {
connection_id: "c1".to_string(),
database: "main".to_string(),
schema: Some("main".to_string()),
object_kinds: vec![CompletionAssistantObjectKind::Table, CompletionAssistantObjectKind::View],
mask: "account".to_string(),
case_sensitive: false,
global_search: false,
max_results: Some(1),
search_in_comments: false,
search_in_definitions: false,
parent_schema: Some("main".to_string()),
parent_name: None,
match_mode: Some(CompletionAssistantMatchMode::Prefix),
},
)
.await
.expect("table completion");
assert_eq!(tables.candidates.len(), 1);
assert!(tables.incomplete);
assert!(!tables.fallback_used);
assert_eq!(tables.candidates[0].name, "account");
let columns = completion_assistant_search(
&pool,
&CompletionAssistantRequest {
connection_id: "c1".to_string(),
database: "main".to_string(),
schema: Some("main".to_string()),
object_kinds: vec![CompletionAssistantObjectKind::Column],
mask: "name".to_string(),
case_sensitive: false,
global_search: false,
max_results: Some(10),
search_in_comments: false,
search_in_definitions: false,
parent_schema: Some("main".to_string()),
parent_name: Some("account".to_string()),
match_mode: Some(CompletionAssistantMatchMode::Contains),
},
)
.await
.expect("column completion");
assert_eq!(columns.candidates.len(), 1);
assert_eq!(columns.candidates[0].name, "display_name");
assert_eq!(columns.candidates[0].data_type.as_deref(), Some("TEXT"));
}
}
pub async fn list_databases(_pool: &SqliteHandle) -> Result<Vec<DatabaseInfo>, String> {
@ -640,6 +706,240 @@ pub async fn get_columns(pool: &SqliteHandle, _schema: &str, table: &str) -> Res
.map_err(|e| e.to_string())?
}
pub async fn completion_assistant_search(
pool: &SqliteHandle,
request: &CompletionAssistantRequest,
) -> Result<CompletionAssistantResponse, String> {
let pool = pool.clone();
let request = request.clone();
tokio::task::spawn_blocking(move || pool.with_connection(|conn| sqlite_completion_assistant_search(conn, &request)))
.await
.map_err(|e| e.to_string())?
}
fn sqlite_completion_assistant_search(
conn: &mut Connection,
request: &CompletionAssistantRequest,
) -> Result<CompletionAssistantResponse, String> {
let limit = request.max_results.unwrap_or(100).clamp(1, 1000);
let kinds = completion_object_kinds(request);
let mut candidates = Vec::new();
if kinds.iter().any(|kind| matches!(kind, CompletionAssistantObjectKind::Schema)) {
for schema in sqlite_completion_schemas(conn, request, limit - candidates.len())? {
candidates.push(schema);
if candidates.len() >= limit {
return Ok(CompletionAssistantResponse { candidates, incomplete: true, fallback_used: false });
}
}
}
if kinds.iter().any(CompletionAssistantObjectKind::is_table_like) {
for table in sqlite_completion_tables(conn, request, &kinds, limit - candidates.len())? {
candidates.push(table);
if candidates.len() >= limit {
return Ok(CompletionAssistantResponse { candidates, incomplete: true, fallback_used: false });
}
}
}
if kinds.iter().any(|kind| matches!(kind, CompletionAssistantObjectKind::Column)) {
for column in sqlite_completion_columns(conn, request, limit - candidates.len())? {
candidates.push(column);
if candidates.len() >= limit {
return Ok(CompletionAssistantResponse { candidates, incomplete: true, fallback_used: false });
}
}
}
Ok(CompletionAssistantResponse { candidates, incomplete: false, fallback_used: false })
}
fn completion_object_kinds(request: &CompletionAssistantRequest) -> Vec<CompletionAssistantObjectKind> {
if request.object_kinds.is_empty() {
vec![CompletionAssistantObjectKind::Table, CompletionAssistantObjectKind::View]
} else {
request.object_kinds.clone()
}
}
fn sqlite_completion_schemas(
conn: &mut Connection,
request: &CompletionAssistantRequest,
limit: usize,
) -> Result<Vec<CompletionAssistantCandidate>, String> {
if limit == 0 {
return Ok(Vec::new());
}
let mut stmt = conn.prepare("PRAGMA database_list").map_err(|e| e.to_string())?;
let rows = stmt.query_map([], |row| row.get::<_, String>(1)).map_err(|e| e.to_string())?;
let mut schemas = rows.collect::<Result<Vec<_>, _>>().map_err(|e| e.to_string())?;
schemas.sort_by_key(|schema| schema.to_lowercase());
Ok(schemas
.into_iter()
.filter(|schema| sqlite_completion_name_matches(schema, request))
.take(limit)
.map(|schema| CompletionAssistantCandidate {
name: schema.clone(),
kind: CompletionAssistantCandidateKind::Schema,
database: Some(request.database.clone()),
schema: Some(schema),
parent_schema: None,
parent_name: None,
comment: None,
data_type: None,
})
.collect())
}
fn sqlite_completion_tables(
conn: &mut Connection,
request: &CompletionAssistantRequest,
kinds: &[CompletionAssistantObjectKind],
limit: usize,
) -> Result<Vec<CompletionAssistantCandidate>, String> {
if limit == 0 {
return Ok(Vec::new());
}
let schema = sqlite_completion_schema(request);
let mut type_filters = Vec::new();
if kinds.iter().any(|kind| matches!(kind, CompletionAssistantObjectKind::Table)) {
type_filters.push("table");
}
if kinds.iter().any(|kind| matches!(kind, CompletionAssistantObjectKind::View)) {
type_filters.push("view");
}
if type_filters.is_empty() {
type_filters.extend(["table", "view"]);
}
let placeholders = std::iter::repeat("?").take(type_filters.len()).collect::<Vec<_>>().join(", ");
let sql = format!(
"SELECT name, type FROM {}.sqlite_master WHERE type IN ({}) AND name NOT LIKE 'sqlite_%' AND {} ORDER BY name LIMIT ?",
sqlite_quote_ident(&schema),
placeholders,
sqlite_completion_filter_sql("name", request)
);
let pattern = sqlite_completion_like_pattern(request);
let mut params: Vec<&dyn rusqlite::ToSql> =
type_filters.iter().map(|value| value as &dyn rusqlite::ToSql).collect();
params.push(&pattern);
params.push(&limit);
let mut stmt = conn.prepare(&sql).map_err(|e| e.to_string())?;
let rows = stmt
.query_map(params.as_slice(), |row| {
let object_type = row.get::<_, String>(1)?;
Ok(CompletionAssistantCandidate {
name: row.get(0)?,
kind: if object_type.eq_ignore_ascii_case("view") {
CompletionAssistantCandidateKind::View
} else {
CompletionAssistantCandidateKind::Table
},
database: Some(request.database.clone()),
schema: Some(schema.clone()),
parent_schema: None,
parent_name: None,
comment: None,
data_type: None,
})
})
.map_err(|e| e.to_string())?;
rows.collect::<Result<Vec<_>, _>>().map_err(|e| e.to_string())
}
fn sqlite_completion_columns(
conn: &mut Connection,
request: &CompletionAssistantRequest,
limit: usize,
) -> Result<Vec<CompletionAssistantCandidate>, String> {
if limit == 0 {
return Ok(Vec::new());
}
let Some(table) = request.parent_name.as_deref().filter(|table| !table.trim().is_empty()) else {
return Ok(Vec::new());
};
let schema = sqlite_completion_schema(request);
let sql = format!("PRAGMA {}.table_info({})", sqlite_quote_ident(&schema), sqlite_quote_string(table));
let mut stmt = conn.prepare(&sql).map_err(|e| e.to_string())?;
let rows = stmt
.query_map([], |row| Ok((row.get::<_, String>("name")?, row.get::<_, String>("type")?)))
.map_err(|e| e.to_string())?;
let mut candidates = Vec::new();
for row in rows {
let (name, data_type) = row.map_err(|e| e.to_string())?;
if !sqlite_completion_name_matches(&name, request) {
continue;
}
candidates.push(CompletionAssistantCandidate {
name,
kind: CompletionAssistantCandidateKind::Column,
database: Some(request.database.clone()),
schema: Some(schema.clone()),
parent_schema: Some(schema.clone()),
parent_name: Some(table.to_string()),
comment: None,
data_type: Some(data_type),
});
if candidates.len() >= limit {
break;
}
}
Ok(candidates)
}
fn sqlite_completion_schema(request: &CompletionAssistantRequest) -> String {
request
.parent_schema
.as_deref()
.or(request.schema.as_deref())
.filter(|schema| !schema.trim().is_empty())
.unwrap_or("main")
.to_string()
}
fn sqlite_completion_name_matches(name: &str, request: &CompletionAssistantRequest) -> bool {
let mask = request.mask.trim().trim_matches('%');
if mask.is_empty() {
return true;
}
let (name, mask) = if request.case_sensitive {
(name.to_string(), mask.to_string())
} else {
(name.to_lowercase(), mask.to_lowercase())
};
match request.match_mode.as_ref().unwrap_or(&CompletionAssistantMatchMode::Prefix) {
CompletionAssistantMatchMode::Prefix => name.starts_with(&mask),
CompletionAssistantMatchMode::Contains => name.contains(&mask),
}
}
fn sqlite_completion_filter_sql(column: &str, request: &CompletionAssistantRequest) -> String {
if request.case_sensitive {
format!("{column} GLOB ?")
} else {
format!("LOWER({column}) LIKE LOWER(?) ESCAPE '\\'")
}
}
fn sqlite_completion_like_pattern(request: &CompletionAssistantRequest) -> String {
let mask = request.mask.trim().trim_matches('%');
let escaped = mask.replace('\\', "\\\\").replace('%', "\\%").replace('_', "\\_");
match request.match_mode.as_ref().unwrap_or(&CompletionAssistantMatchMode::Prefix) {
CompletionAssistantMatchMode::Prefix if request.case_sensitive => format!("{}*", mask.replace('[', "[[]")),
CompletionAssistantMatchMode::Contains if request.case_sensitive => format!("*{}*", mask.replace('[', "[[]")),
CompletionAssistantMatchMode::Prefix => format!("{escaped}%"),
CompletionAssistantMatchMode::Contains => format!("%{escaped}%"),
}
}
fn sqlite_quote_ident(value: &str) -> String {
format!("\"{}\"", value.replace('"', "\"\""))
}
fn sqlite_quote_string(value: &str) -> String {
format!("'{}'", value.replace('\'', "''"))
}
/// Read `sqlite_master.sql` for `table` and return the lowercase column names that
/// are rowid-alias autoincrement primary keys (i.e. SQLite will assign a value when
/// the column is omitted from an INSERT). Returns `None` only on connection / query

View File

@ -842,9 +842,215 @@ pub async fn list_tables(
limit: Option<usize>,
offset: Option<usize>,
) -> Result<Vec<TableInfo>, String> {
let sql = sqlserver_list_tables_sql(schema, filter, limit, offset);
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
Ok(rows
.iter()
.map(|row| TableInfo {
name: row.get::<&str, _>(0).unwrap_or("").to_string(),
table_type: row.get::<&str, _>(1).unwrap_or("BASE TABLE").to_string(),
comment: row.get::<&str, _>(2).filter(|s: &&str| !s.is_empty()).map(|s: &str| s.to_string()),
parent_schema: None,
parent_name: None,
})
.collect())
}
pub async fn completion_assistant_search(
client: &mut SqlServerClient,
request: &crate::types::CompletionAssistantRequest,
) -> Result<crate::types::CompletionAssistantResponse, String> {
let limit = request.max_results.unwrap_or(100).clamp(1, 1000);
let sql = sqlserver_completion_assistant_sql(request, limit);
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
let candidates = rows
.iter()
.map(|row| {
let object_type = row.get::<&str, _>(2).unwrap_or("OBJECT");
crate::types::CompletionAssistantCandidate {
name: row.get::<&str, _>(0).unwrap_or("").to_string(),
kind: sqlserver_completion_candidate_kind(object_type),
database: Some(request.database.clone()),
schema: row.get::<&str, _>(1).map(str::to_string),
parent_schema: row.get::<&str, _>(3).map(str::to_string),
parent_name: row.get::<&str, _>(4).map(str::to_string),
comment: row.get::<&str, _>(5).filter(|s: &&str| !s.is_empty()).map(|s| (*s).to_string()),
data_type: row.get::<&str, _>(6).map(str::to_string),
}
})
.collect::<Vec<_>>();
Ok(crate::types::CompletionAssistantResponse {
incomplete: candidates.len() >= limit,
candidates,
fallback_used: false,
})
}
fn sqlserver_completion_candidate_kind(object_type: &str) -> crate::types::CompletionAssistantCandidateKind {
match object_type.to_ascii_uppercase().as_str() {
"SCHEMA" => crate::types::CompletionAssistantCandidateKind::Schema,
"TABLE" | "BASE TABLE" => crate::types::CompletionAssistantCandidateKind::Table,
"VIEW" => crate::types::CompletionAssistantCandidateKind::View,
"PROCEDURE" => crate::types::CompletionAssistantCandidateKind::Procedure,
"FUNCTION" => crate::types::CompletionAssistantCandidateKind::Function,
"COLUMN" => crate::types::CompletionAssistantCandidateKind::Column,
_ => crate::types::CompletionAssistantCandidateKind::Object,
}
}
fn sqlserver_completion_assistant_sql(request: &crate::types::CompletionAssistantRequest, limit: usize) -> String {
let object_kinds = if request.object_kinds.is_empty() {
vec![crate::types::CompletionAssistantObjectKind::Table, crate::types::CompletionAssistantObjectKind::View]
} else {
request.object_kinds.clone()
};
let mask = request.mask.trim();
let like_pattern = completion_like_pattern(mask, request.match_mode.as_ref());
let like_clause = if like_pattern == "%" {
String::new()
} else {
format!(" AND LOWER({}) LIKE LOWER('{like_pattern}') ESCAPE '\\' ", "name_expr")
};
let schema_filter = request
.schema
.as_deref()
.or(request.parent_schema.as_deref())
.filter(|schema| !schema.trim().is_empty())
.map(|schema| format!(" AND s.name = '{}' ", schema.replace('\'', "''")))
.unwrap_or_default();
let mut queries = Vec::new();
if (mask.starts_with('#') || mask.starts_with("%#"))
&& object_kinds.iter().any(crate::types::CompletionAssistantObjectKind::is_table_like)
{
let object_like = sqlserver_completion_object_search_clause(request, &like_pattern);
queries.push(format!(
"SELECT TOP ({limit}) o.name, s.name AS schema_name, 'TABLE' AS object_type, NULL AS parent_schema, NULL AS parent_name, NULL AS object_comment, NULL AS data_type \
FROM tempdb.sys.all_objects o \
JOIN tempdb.sys.schemas s ON s.schema_id = o.schema_id \
WHERE o.type = 'U' {object_like}"
));
return format!("SELECT * FROM ({}) AS dbx_completion ORDER BY name", queries.remove(0));
}
if object_kinds.iter().any(|kind| matches!(kind, crate::types::CompletionAssistantObjectKind::Schema)) {
let schema_like = like_clause.replace("name_expr", "s.name");
queries.push(format!(
"SELECT TOP ({limit}) s.name, s.name AS schema_name, 'SCHEMA' AS object_type, NULL AS parent_schema, NULL AS parent_name, NULL AS object_comment, NULL AS data_type \
FROM sys.schemas s \
WHERE s.name NOT IN ('guest','INFORMATION_SCHEMA','sys') {schema_like}"
));
}
if object_kinds.iter().any(crate::types::CompletionAssistantObjectKind::is_table_like)
|| object_kinds.iter().any(crate::types::CompletionAssistantObjectKind::is_routine_like)
{
let mut type_ids = Vec::new();
if object_kinds.iter().any(|kind| matches!(kind, crate::types::CompletionAssistantObjectKind::Table)) {
type_ids.push("'U'");
}
if object_kinds.iter().any(|kind| matches!(kind, crate::types::CompletionAssistantObjectKind::View)) {
type_ids.push("'V'");
}
if object_kinds.iter().any(|kind| {
matches!(
kind,
crate::types::CompletionAssistantObjectKind::Procedure
| crate::types::CompletionAssistantObjectKind::Routine
)
}) {
type_ids.push("'P'");
}
if object_kinds.iter().any(|kind| {
matches!(
kind,
crate::types::CompletionAssistantObjectKind::Function
| crate::types::CompletionAssistantObjectKind::Routine
)
}) {
type_ids.extend(["'FN'", "'IF'", "'TF'", "'FS'", "'FT'"]);
}
let object_like = sqlserver_completion_object_search_clause(request, &like_pattern);
queries.push(format!(
"SELECT TOP ({limit}) o.name, s.name AS schema_name, \
CASE o.type WHEN 'U' THEN 'TABLE' WHEN 'V' THEN 'VIEW' WHEN 'P' THEN 'PROCEDURE' WHEN 'FN' THEN 'FUNCTION' WHEN 'IF' THEN 'FUNCTION' WHEN 'TF' THEN 'FUNCTION' WHEN 'FS' THEN 'FUNCTION' WHEN 'FT' THEN 'FUNCTION' ELSE o.type_desc END AS object_type, \
NULL AS parent_schema, NULL AS parent_name, ep.value AS object_comment, NULL AS data_type \
FROM sys.objects o \
JOIN sys.schemas s ON s.schema_id = o.schema_id \
OUTER APPLY (SELECT CAST(ep.value AS NVARCHAR(MAX)) AS value FROM sys.extended_properties ep WHERE ep.major_id = o.object_id AND ep.minor_id = 0 AND ep.name = N'MS_Description') ep \
WHERE o.type IN ({}) AND o.is_ms_shipped = 0 {schema_filter} {object_like}",
type_ids.join(",")
));
}
if object_kinds.iter().any(|kind| matches!(kind, crate::types::CompletionAssistantObjectKind::Column)) {
let column_like = like_clause.replace("name_expr", "c.name");
let parent_table_filter = request
.parent_name
.as_deref()
.filter(|table| !table.trim().is_empty())
.map(|table| format!(" AND o.name = '{}' ", table.replace('\'', "''")))
.unwrap_or_default();
queries.push(format!(
"SELECT TOP ({limit}) c.name, s.name AS schema_name, 'COLUMN' AS object_type, s.name AS parent_schema, o.name AS parent_name, NULL AS object_comment, TYPE_NAME(c.user_type_id) AS data_type \
FROM sys.columns c \
JOIN sys.objects o ON o.object_id = c.object_id \
JOIN sys.schemas s ON s.schema_id = o.schema_id \
WHERE o.type IN ('U','V') AND o.is_ms_shipped = 0 {schema_filter} {parent_table_filter} {column_like}"
));
}
if queries.is_empty() {
format!("SELECT TOP (0) '' AS name, '' AS schema_name, '' AS object_type, NULL AS parent_schema, NULL AS parent_name, NULL AS object_comment, NULL AS data_type")
} else if queries.len() == 1 {
format!("SELECT * FROM ({}) AS dbx_completion ORDER BY name", queries.remove(0))
} else {
format!("SELECT TOP ({limit}) * FROM ({}) AS dbx_completion ORDER BY name", queries.join(" UNION ALL "))
}
}
fn sqlserver_completion_object_search_clause(
request: &crate::types::CompletionAssistantRequest,
like_pattern: &str,
) -> String {
if like_pattern == "%" {
return String::new();
}
let mut predicates = vec![format!("LOWER(o.name) LIKE LOWER('{like_pattern}') ESCAPE '\\'")];
if request.search_in_comments {
predicates.push(format!("LOWER(COALESCE(ep.value, '')) LIKE LOWER('{like_pattern}') ESCAPE '\\'"));
}
if request.search_in_definitions {
predicates.push(format!(
"LOWER(COALESCE(OBJECT_DEFINITION(o.object_id), '')) LIKE LOWER('{like_pattern}') ESCAPE '\\'"
));
}
format!(" AND ({}) ", predicates.join(" OR "))
}
fn completion_like_pattern(mask: &str, mode: Option<&crate::types::CompletionAssistantMatchMode>) -> String {
if mask.is_empty() || mask == "%" {
return "%".to_string();
}
let has_wildcard = mask.contains('%');
if has_wildcard {
return mask.split('%').map(escape_like_literal).collect::<Vec<_>>().join("%");
}
let escaped = escape_like_literal(mask);
match mode.unwrap_or(&crate::types::CompletionAssistantMatchMode::Prefix) {
crate::types::CompletionAssistantMatchMode::Prefix => format!("{escaped}%"),
crate::types::CompletionAssistantMatchMode::Contains => format!("%{escaped}%"),
}
}
fn sqlserver_list_tables_sql(
schema: &str,
filter: Option<&str>,
limit: Option<usize>,
offset: Option<usize>,
) -> String {
let filter_clause = filter
.filter(|value| !value.trim().is_empty())
.map(|value| format!(" AND o.name LIKE '%{}%' ESCAPE '\\' ", escape_like_literal(value.trim())))
.map(|value| format!(" AND LOWER(o.name) LIKE LOWER('%{}%') ESCAPE '\\' ", escape_like_literal(value.trim())))
.unwrap_or_default();
let schema_escaped = schema.replace('\'', "''");
let base_columns = "o.name, CASE WHEN o.type = 'V' THEN 'VIEW' ELSE 'BASE TABLE' END, ep.value AS TABLE_COMMENT";
@ -858,7 +1064,7 @@ pub async fn list_tables(
// Use SELECT TOP for broad SQL Server version compatibility.
// OFFSET / FETCH NEXT is only available in SQL Server 2012+.
let sql = match (limit, offset) {
match (limit, offset) {
(Some(limit), Some(offset)) if offset > 0 => {
let end = offset + limit.min(1000);
format!(
@ -874,19 +1080,7 @@ pub async fn list_tables(
_ => {
format!("SELECT {base_columns} {base_from} {base_where} {order_by}")
}
};
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
Ok(rows
.iter()
.map(|row| TableInfo {
name: row.get::<&str, _>(0).unwrap_or("").to_string(),
table_type: row.get::<&str, _>(1).unwrap_or("BASE TABLE").to_string(),
comment: row.get::<&str, _>(2).filter(|s: &&str| !s.is_empty()).map(|s: &str| s.to_string()),
parent_schema: None,
parent_name: None,
})
.collect())
}
}
fn escape_like_literal(value: &str) -> String {
@ -1414,9 +1608,11 @@ mod tests {
use super::{
build_spatial_safe_sqlserver_query, is_sqlserver_spatial_column, requires_simple_query_batch,
sqlserver_batch_can_use_execute, sqlserver_cell_to_json, sqlserver_columns_sql,
sqlserver_dml_output_returns_rows, sqlserver_indexes_sql, sqlserver_list_objects_sql,
sqlserver_table_comment_sql, SqlServerDescribedColumn, SqlServerResultSet,
sqlserver_completion_assistant_sql, sqlserver_dml_output_returns_rows, sqlserver_indexes_sql,
sqlserver_list_objects_sql, sqlserver_list_tables_sql, sqlserver_table_comment_sql, SqlServerDescribedColumn,
SqlServerResultSet,
};
use crate::types::{CompletionAssistantMatchMode, CompletionAssistantObjectKind, CompletionAssistantRequest};
use chrono::NaiveDate;
use std::time::Instant;
use tiberius::{ColumnData, IntoSql};
@ -1556,6 +1752,137 @@ mod tests {
assert!(sql.contains("modify_date"));
}
#[test]
fn sqlserver_list_tables_filter_is_case_insensitive() {
let sql = sqlserver_list_tables_sql("dbo", Some("temp"), Some(200), None);
assert!(sql.contains("LOWER(o.name) LIKE LOWER('%temp%') ESCAPE '\\'"));
assert!(sql.contains("SELECT TOP (200)"));
}
#[test]
fn sqlserver_list_tables_filter_escapes_like_literals() {
let sql = sqlserver_list_tables_sql("dbo", Some("Temp_Table[%]"), Some(200), None);
assert!(sql.contains("LOWER(o.name) LIKE LOWER('%Temp\\_Table\\[\\%]%') ESCAPE '\\'"));
}
#[test]
fn sqlserver_completion_assistant_searches_objects_before_limiting() {
let request = CompletionAssistantRequest {
connection_id: "c1".to_string(),
database: "app".to_string(),
schema: Some("dbo".to_string()),
object_kinds: vec![CompletionAssistantObjectKind::Table, CompletionAssistantObjectKind::View],
mask: "Temp".to_string(),
case_sensitive: false,
global_search: false,
max_results: Some(100),
search_in_comments: false,
search_in_definitions: false,
parent_schema: None,
parent_name: None,
match_mode: Some(CompletionAssistantMatchMode::Prefix),
};
let sql = sqlserver_completion_assistant_sql(&request, 100);
assert!(sql.contains("SELECT TOP (100)"));
assert!(sql.contains("FROM sys.objects o"));
assert!(sql.contains("o.type IN ('U','V')"));
assert!(sql.contains("s.name = 'dbo'"));
assert!(sql.contains("LOWER(o.name) LIKE LOWER('Temp%') ESCAPE '\\'"));
}
#[test]
fn sqlserver_completion_assistant_searches_columns_by_parent_table() {
let request = CompletionAssistantRequest {
connection_id: "c1".to_string(),
database: "app".to_string(),
schema: Some("dbo".to_string()),
object_kinds: vec![CompletionAssistantObjectKind::Column],
mask: "id".to_string(),
case_sensitive: false,
global_search: false,
max_results: Some(50),
search_in_comments: false,
search_in_definitions: false,
parent_schema: Some("dbo".to_string()),
parent_name: Some("Users".to_string()),
match_mode: Some(CompletionAssistantMatchMode::Contains),
};
let sql = sqlserver_completion_assistant_sql(&request, 50);
assert!(sql.contains("FROM sys.columns c"));
assert!(sql.contains("o.name = 'Users'"));
assert!(sql.contains("LOWER(c.name) LIKE LOWER('%id%') ESCAPE '\\'"));
}
#[test]
fn sqlserver_completion_assistant_searches_tempdb_for_temp_table_masks() {
let request = CompletionAssistantRequest {
connection_id: "c1".to_string(),
database: "app".to_string(),
schema: Some("dbo".to_string()),
object_kinds: vec![CompletionAssistantObjectKind::Table],
mask: "#Temp".to_string(),
case_sensitive: false,
global_search: false,
max_results: Some(100),
search_in_comments: false,
search_in_definitions: false,
parent_schema: None,
parent_name: None,
match_mode: Some(CompletionAssistantMatchMode::Prefix),
};
let sql = sqlserver_completion_assistant_sql(&request, 100);
assert!(sql.contains("FROM tempdb.sys.all_objects o"));
assert!(sql.contains("o.type = 'U'"));
assert!(sql.contains("LOWER(o.name) LIKE LOWER('#Temp%') ESCAPE '\\'"));
}
#[test]
fn sqlserver_completion_assistant_generates_scoped_search_masks() {
assert_eq!(super::completion_like_pattern("Temp", Some(&CompletionAssistantMatchMode::Prefix)), "Temp%");
assert_eq!(super::completion_like_pattern("Temp", Some(&CompletionAssistantMatchMode::Contains)), "%Temp%");
assert_eq!(
super::completion_like_pattern("dbo.Temp%", Some(&CompletionAssistantMatchMode::Prefix)),
"dbo.Temp%"
);
assert_eq!(
super::completion_like_pattern("Temp_Table", Some(&CompletionAssistantMatchMode::Prefix)),
"Temp\\_Table%"
);
}
#[test]
fn sqlserver_completion_assistant_can_search_comments_and_definitions() {
let request = CompletionAssistantRequest {
connection_id: "c1".to_string(),
database: "app".to_string(),
schema: Some("dbo".to_string()),
object_kinds: vec![CompletionAssistantObjectKind::Procedure],
mask: "audit".to_string(),
case_sensitive: false,
global_search: false,
max_results: Some(100),
search_in_comments: true,
search_in_definitions: true,
parent_schema: None,
parent_name: None,
match_mode: Some(CompletionAssistantMatchMode::Contains),
};
let sql = sqlserver_completion_assistant_sql(&request, 100);
assert!(sql.contains("COALESCE(ep.value, '')"));
assert!(sql.contains("OBJECT_DEFINITION(o.object_id)"));
assert!(sql.contains("LOWER('%audit%')"));
}
#[test]
fn sqlserver_tinyint_cells_are_json_numbers() {
assert_eq!(sqlserver_cell_to_json(&ColumnData::U8(Some(7))), serde_json::json!(7));

View File

@ -4,7 +4,7 @@ use crate::models::connection::{ConnectionConfig, DatabaseType};
use crate::query::{agent_execute_query_params, should_discard_pool_after_error, QueryExecutionOptions};
use std::future::Future;
use std::sync::Arc;
use std::time::Duration;
use std::time::{Duration, Instant};
macro_rules! extract_pool {
($connections:expr, $key:expr, $variant:ident) => {
@ -239,6 +239,190 @@ pub fn duckdb_query_columns_in_database_with_attached(
Ok(rows.filter_map(|r| r.ok()).collect())
}
#[cfg(feature = "duckdb-bundled")]
pub fn duckdb_completion_assistant_search(
con: &duckdb::Connection,
request: &db::CompletionAssistantRequest,
attached_names: &[String],
) -> Result<db::CompletionAssistantResponse, String> {
let limit = request.max_results.unwrap_or(100).clamp(1, 1000);
let kinds = if request.object_kinds.is_empty() {
vec![db::CompletionAssistantObjectKind::Table, db::CompletionAssistantObjectKind::View]
} else {
request.object_kinds.clone()
};
let mut candidates = Vec::new();
if kinds.iter().any(|kind| matches!(kind, db::CompletionAssistantObjectKind::Schema)) {
candidates.extend(duckdb_completion_schemas(con, request, attached_names, limit)?);
if candidates.len() >= limit {
return Ok(db::CompletionAssistantResponse { candidates, incomplete: true, fallback_used: false });
}
}
if kinds.iter().any(db::CompletionAssistantObjectKind::is_table_like) {
candidates.extend(duckdb_completion_tables(con, request, &kinds, attached_names, limit - candidates.len())?);
if candidates.len() >= limit {
return Ok(db::CompletionAssistantResponse { candidates, incomplete: true, fallback_used: false });
}
}
if kinds.iter().any(|kind| matches!(kind, db::CompletionAssistantObjectKind::Column)) {
candidates.extend(duckdb_completion_columns(con, request, attached_names, limit - candidates.len())?);
if candidates.len() >= limit {
return Ok(db::CompletionAssistantResponse { candidates, incomplete: true, fallback_used: false });
}
}
Ok(db::CompletionAssistantResponse { candidates, incomplete: false, fallback_used: false })
}
#[cfg(feature = "duckdb-bundled")]
fn duckdb_completion_schemas(
con: &duckdb::Connection,
request: &db::CompletionAssistantRequest,
attached_names: &[String],
limit: usize,
) -> Result<Vec<db::CompletionAssistantCandidate>, String> {
if limit == 0 {
return Ok(Vec::new());
}
let database = duckdb_catalog_name(con, &request.database, attached_names)?;
let pattern = duckdb_completion_like_pattern(request);
let mut stmt = con
.prepare(
"SELECT schema_name
FROM information_schema.schemata
WHERE catalog_name = ?
AND schema_name NOT IN ('information_schema', 'pg_catalog')
AND lower(schema_name) LIKE lower(?) ESCAPE '\\'
ORDER BY schema_name
LIMIT ?",
)
.map_err(|e| e.to_string())?;
let rows = stmt
.query_map((database.as_str(), pattern.as_str(), limit as i64), |row| {
let schema = row.get::<_, String>(0)?;
Ok(db::CompletionAssistantCandidate {
name: schema.clone(),
kind: db::CompletionAssistantCandidateKind::Schema,
database: Some(request.database.clone()),
schema: Some(schema),
parent_schema: None,
parent_name: None,
comment: None,
data_type: None,
})
})
.map_err(|e| e.to_string())?;
Ok(rows.filter_map(|row| row.ok()).collect())
}
#[cfg(feature = "duckdb-bundled")]
fn duckdb_completion_tables(
con: &duckdb::Connection,
request: &db::CompletionAssistantRequest,
kinds: &[db::CompletionAssistantObjectKind],
attached_names: &[String],
limit: usize,
) -> Result<Vec<db::CompletionAssistantCandidate>, String> {
if limit == 0 {
return Ok(Vec::new());
}
let database = duckdb_catalog_name(con, &request.database, attached_names)?;
let schema = request.parent_schema.as_deref().or(request.schema.as_deref()).unwrap_or("main");
let include_tables = kinds.iter().any(|kind| matches!(kind, db::CompletionAssistantObjectKind::Table));
let include_views = kinds.iter().any(|kind| matches!(kind, db::CompletionAssistantObjectKind::View));
let pattern = duckdb_completion_like_pattern(request);
let mut stmt = con
.prepare(
"SELECT table_name, table_type
FROM information_schema.tables
WHERE table_catalog = ?
AND table_schema = ?
AND ((? AND table_type = 'BASE TABLE') OR (? AND table_type = 'VIEW'))
AND lower(table_name) LIKE lower(?) ESCAPE '\\'
ORDER BY table_name
LIMIT ?",
)
.map_err(|e| e.to_string())?;
let rows = stmt
.query_map((database.as_str(), schema, include_tables, include_views, pattern.as_str(), limit as i64), |row| {
let table_type = row.get::<_, String>(1)?;
Ok(db::CompletionAssistantCandidate {
name: row.get(0)?,
kind: if table_type.eq_ignore_ascii_case("VIEW") {
db::CompletionAssistantCandidateKind::View
} else {
db::CompletionAssistantCandidateKind::Table
},
database: Some(request.database.clone()),
schema: Some(schema.to_string()),
parent_schema: None,
parent_name: None,
comment: None,
data_type: None,
})
})
.map_err(|e| e.to_string())?;
Ok(rows.filter_map(|row| row.ok()).collect())
}
#[cfg(feature = "duckdb-bundled")]
fn duckdb_completion_columns(
con: &duckdb::Connection,
request: &db::CompletionAssistantRequest,
attached_names: &[String],
limit: usize,
) -> Result<Vec<db::CompletionAssistantCandidate>, String> {
if limit == 0 {
return Ok(Vec::new());
}
let Some(table) = request.parent_name.as_deref().filter(|table| !table.trim().is_empty()) else {
return Ok(Vec::new());
};
let database = duckdb_catalog_name(con, &request.database, attached_names)?;
let schema = request.parent_schema.as_deref().or(request.schema.as_deref()).unwrap_or("main");
let pattern = duckdb_completion_like_pattern(request);
let mut stmt = con
.prepare(
"SELECT column_name, data_type
FROM information_schema.columns
WHERE table_catalog = ?
AND table_schema = ?
AND table_name = ?
AND lower(column_name) LIKE lower(?) ESCAPE '\\'
ORDER BY ordinal_position
LIMIT ?",
)
.map_err(|e| e.to_string())?;
let rows = stmt
.query_map((database.as_str(), schema, table, pattern.as_str(), limit as i64), |row| {
Ok(db::CompletionAssistantCandidate {
name: row.get(0)?,
kind: db::CompletionAssistantCandidateKind::Column,
database: Some(request.database.clone()),
schema: Some(schema.to_string()),
parent_schema: Some(schema.to_string()),
parent_name: Some(table.to_string()),
comment: None,
data_type: Some(row.get(1)?),
})
})
.map_err(|e| e.to_string())?;
Ok(rows.filter_map(|row| row.ok()).collect())
}
#[cfg(feature = "duckdb-bundled")]
fn duckdb_completion_like_pattern(request: &db::CompletionAssistantRequest) -> String {
let mask = request.mask.trim().trim_matches('%');
let escaped = mask.replace('\\', "\\\\").replace('%', "\\%").replace('_', "\\_");
match request.match_mode.as_ref().unwrap_or(&db::CompletionAssistantMatchMode::Prefix) {
db::CompletionAssistantMatchMode::Prefix => format!("{escaped}%"),
db::CompletionAssistantMatchMode::Contains => format!("%{escaped}%"),
}
}
#[cfg(feature = "duckdb-bundled")]
async fn duckdb_attached_database_names(state: &AppState, connection_id: &str) -> Vec<String> {
state
@ -1164,6 +1348,65 @@ mod tests {
let _ = std::fs::remove_file(path);
}
#[cfg(feature = "duckdb-bundled")]
#[test]
fn duckdb_completion_assistant_searches_catalog_metadata_with_limit() {
let con = duckdb::Connection::open_in_memory().unwrap();
con.execute_batch(
"CREATE TABLE account(id INTEGER, display_name VARCHAR); CREATE VIEW account_view AS SELECT id FROM account;",
)
.unwrap();
let request = db::CompletionAssistantRequest {
connection_id: "c1".to_string(),
database: "main".to_string(),
schema: Some("main".to_string()),
object_kinds: vec![db::CompletionAssistantObjectKind::Table, db::CompletionAssistantObjectKind::View],
mask: "account".to_string(),
case_sensitive: false,
global_search: false,
max_results: Some(1),
search_in_comments: false,
search_in_definitions: false,
parent_schema: Some("main".to_string()),
parent_name: None,
match_mode: Some(db::CompletionAssistantMatchMode::Prefix),
};
let tables = duckdb_completion_assistant_search(&con, &request, &[]).unwrap();
assert_eq!(tables.candidates.len(), 1);
assert!(tables.incomplete);
assert!(!tables.fallback_used);
assert_eq!(tables.candidates[0].name, "account");
let columns = duckdb_completion_assistant_search(
&con,
&db::CompletionAssistantRequest {
object_kinds: vec![db::CompletionAssistantObjectKind::Column],
mask: "name".to_string(),
max_results: Some(10),
parent_name: Some("account".to_string()),
match_mode: Some(db::CompletionAssistantMatchMode::Contains),
..request
},
&[],
)
.unwrap();
assert_eq!(columns.candidates.len(), 1);
assert_eq!(columns.candidates[0].name, "display_name");
}
#[test]
fn detects_unsupported_agent_completion_assistant_errors() {
assert!(super::is_agent_completion_assistant_unsupported(
"Agent RPC error (-1): Unknown method: completion_assistant_search_v1"
));
assert!(super::is_agent_completion_assistant_unsupported(
"Agent RPC error (-1): Completion assistant search is not supported by this agent"
));
assert!(!super::is_agent_completion_assistant_unsupported("Agent RPC error (-1): Connection failed"));
}
#[test]
fn clickhouse_metadata_uses_schema_when_database_is_empty() {
assert_eq!(clickhouse_metadata_database("", "testdb"), "testdb");
@ -1248,6 +1491,248 @@ pub async fn list_completion_objects_core(
.await
}
pub async fn completion_assistant_search_core(
state: &AppState,
request: db::CompletionAssistantRequest,
) -> Result<db::CompletionAssistantResponse, String> {
let started_at = Instant::now();
let request_summary = format!(
"connection_id={} database={} schema={:?} kinds={:?} mask={} limit={:?}",
request.connection_id,
request.database,
request.schema,
request.object_kinds,
request.mask,
request.max_results
);
retry_metadata_connection(state, &request.connection_id, Some(&request.database), || async {
let pool_key = state.get_or_create_pool(&request.connection_id, Some(&request.database)).await?;
log::debug!("[schema][completion_assistant:start] {request_summary}");
{
let connections = state.connections.read().await;
try_sqlserver!(connections, &pool_key, completion_assistant_search, &request);
}
{
let connections = state.connections.read().await;
if let Some(pool) = connections.get(&pool_key).and_then(|pool| match pool {
PoolKind::Sqlite(pool) => Some(pool.clone()),
_ => None,
}) {
drop(connections);
return db::sqlite::completion_assistant_search(&pool, &request).await;
}
}
#[cfg(feature = "duckdb-bundled")]
{
let duckdb_attached_names = duckdb_attached_database_names(state, &request.connection_id).await;
let connections = state.connections.read().await;
if let Some(con) = extract_pool!(&connections, &pool_key, DuckDb) {
drop(connections);
let con = con.lock().map_err(|e| e.to_string())?;
return duckdb_completion_assistant_search(&con, &request, &duckdb_attached_names);
}
}
{
let connections = state.connections.read().await;
if let Some(pool) = connections.get(&pool_key).and_then(|pool| match pool {
PoolKind::Postgres(pool) => Some(pool.clone()),
_ => None,
}) {
drop(connections);
return db::postgres::completion_assistant_search(&pool, &request).await;
}
}
{
let connections = state.connections.read().await;
if let Some(pool) = connections.get(&pool_key).and_then(|pool| match pool {
PoolKind::Mysql(pool, mode) if *mode != MysqlMode::OceanBaseOracle => Some(pool.clone()),
_ => None,
}) {
drop(connections);
return db::mysql::completion_assistant_search(&pool, &request).await;
}
}
{
let connections = state.connections.read().await;
if let Some(client) = extract_pool!(&connections, &pool_key, Agent) {
let db_config = connection_config(state, &request.connection_id).await;
drop(connections);
let mut client = client.lock().await;
match client
.completion_assistant_search::<db::CompletionAssistantResponse>(
&request,
agent_metadata_timeout(db_config.as_ref()),
)
.await
{
Ok(mut response) => {
response.fallback_used = false;
return Ok(response);
}
Err(error) if is_agent_completion_assistant_unsupported(&error) => {
log::debug!(
"[schema][completion_assistant:agent-fallback] {} reason={}",
request_summary,
error
);
}
Err(error) => return Err(error),
}
}
}
let response = completion_assistant_fallback_core(state, &request).await;
if let Ok(response) = &response {
log::debug!(
"[schema][completion_assistant:done] {} elapsed_ms={} candidates={} fallback_used={}",
request_summary,
started_at.elapsed().as_millis(),
response.candidates.len(),
response.fallback_used
);
}
response
})
.await
}
fn is_agent_completion_assistant_unsupported(error: &str) -> bool {
error.contains("Unknown method: completion_assistant_search_v1")
|| error.contains("Method not found: completion_assistant_search_v1")
|| error.contains("method not found: completion_assistant_search_v1")
|| error.contains("Completion assistant search is not supported")
}
async fn completion_assistant_fallback_core(
state: &AppState,
request: &db::CompletionAssistantRequest,
) -> Result<db::CompletionAssistantResponse, String> {
let limit = request.max_results.unwrap_or(100).clamp(1, 1000);
let kinds = if request.object_kinds.is_empty() {
vec![db::CompletionAssistantObjectKind::Table, db::CompletionAssistantObjectKind::View]
} else {
request.object_kinds.clone()
};
let mut candidates = Vec::new();
let schema = request.parent_schema.as_deref().or(request.schema.as_deref()).unwrap_or("");
let filter = request.mask.trim().trim_matches('%');
if kinds.iter().any(|kind| matches!(kind, db::CompletionAssistantObjectKind::Schema)) {
let schemas = list_schemas_core(state, &request.connection_id, &request.database).await?;
for schema_name in schemas {
if completion_name_matches(&schema_name, filter, request.match_mode.as_ref()) {
candidates.push(db::CompletionAssistantCandidate {
name: schema_name.clone(),
kind: db::CompletionAssistantCandidateKind::Schema,
database: Some(request.database.clone()),
schema: Some(schema_name),
parent_schema: None,
parent_name: None,
comment: None,
data_type: None,
});
}
if candidates.len() >= limit {
return Ok(db::CompletionAssistantResponse { candidates, incomplete: true, fallback_used: true });
}
}
}
if kinds.iter().any(db::CompletionAssistantObjectKind::is_table_like) {
let object_types = completion_table_object_types(&kinds);
let tables = list_tables_core(
state,
&request.connection_id,
&request.database,
schema,
if filter.is_empty() { None } else { Some(filter) },
Some(limit),
None,
object_types.as_deref(),
)
.await?;
for table in tables {
let kind = if table.table_type.eq_ignore_ascii_case("VIEW")
|| table.table_type.eq_ignore_ascii_case("MATERIALIZED_VIEW")
{
db::CompletionAssistantCandidateKind::View
} else {
db::CompletionAssistantCandidateKind::Table
};
candidates.push(db::CompletionAssistantCandidate {
name: table.name,
kind,
database: Some(request.database.clone()),
schema: if schema.is_empty() { None } else { Some(schema.to_string()) },
parent_schema: table.parent_schema,
parent_name: table.parent_name,
comment: table.comment,
data_type: None,
});
if candidates.len() >= limit {
return Ok(db::CompletionAssistantResponse { candidates, incomplete: true, fallback_used: true });
}
}
}
if kinds.iter().any(|kind| matches!(kind, db::CompletionAssistantObjectKind::Column)) {
if let Some(table) = request.parent_name.as_deref().filter(|table| !table.trim().is_empty()) {
let columns = get_columns_core(state, &request.connection_id, &request.database, schema, table).await?;
for column in columns {
if completion_name_matches(&column.name, filter, request.match_mode.as_ref()) {
candidates.push(db::CompletionAssistantCandidate {
name: column.name,
kind: db::CompletionAssistantCandidateKind::Column,
database: Some(request.database.clone()),
schema: if schema.is_empty() { None } else { Some(schema.to_string()) },
parent_schema: if schema.is_empty() { None } else { Some(schema.to_string()) },
parent_name: Some(table.to_string()),
comment: column.comment,
data_type: Some(column.data_type),
});
}
if candidates.len() >= limit {
return Ok(db::CompletionAssistantResponse { candidates, incomplete: true, fallback_used: true });
}
}
}
}
Ok(db::CompletionAssistantResponse { candidates, incomplete: false, fallback_used: true })
}
fn completion_table_object_types(kinds: &[db::CompletionAssistantObjectKind]) -> Option<Vec<String>> {
let mut object_types = Vec::new();
if kinds.iter().any(|kind| matches!(kind, db::CompletionAssistantObjectKind::Table)) {
object_types.push("table".to_string());
}
if kinds.iter().any(|kind| matches!(kind, db::CompletionAssistantObjectKind::View)) {
object_types.push("view".to_string());
}
if object_types.is_empty() {
None
} else {
Some(object_types)
}
}
fn completion_name_matches(name: &str, filter: &str, mode: Option<&db::CompletionAssistantMatchMode>) -> bool {
if filter.is_empty() {
return true;
}
let name = name.to_lowercase();
let filter = filter.to_lowercase();
match mode.unwrap_or(&db::CompletionAssistantMatchMode::Prefix) {
db::CompletionAssistantMatchMode::Prefix => name.starts_with(&filter),
db::CompletionAssistantMatchMode::Contains => name.contains(&filter),
}
}
async fn list_object_statistics_once(
state: &AppState,
connection_id: &str,

View File

@ -75,6 +75,91 @@ pub struct ColumnInfo {
pub character_maximum_length: Option<i32>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum CompletionAssistantObjectKind {
Database,
Schema,
Table,
View,
Routine,
Procedure,
Function,
Column,
}
impl CompletionAssistantObjectKind {
pub fn is_table_like(&self) -> bool {
matches!(self, Self::Table | Self::View)
}
pub fn is_routine_like(&self) -> bool {
matches!(self, Self::Routine | Self::Procedure | Self::Function)
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum CompletionAssistantCandidateKind {
Database,
Schema,
Table,
View,
Procedure,
Function,
Column,
Object,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum CompletionAssistantMatchMode {
Prefix,
Contains,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompletionAssistantRequest {
pub connection_id: String,
pub database: String,
pub schema: Option<String>,
#[serde(default)]
pub object_kinds: Vec<CompletionAssistantObjectKind>,
#[serde(default)]
pub mask: String,
#[serde(default)]
pub case_sensitive: bool,
#[serde(default)]
pub global_search: bool,
pub max_results: Option<usize>,
#[serde(default)]
pub search_in_comments: bool,
#[serde(default)]
pub search_in_definitions: bool,
pub parent_schema: Option<String>,
pub parent_name: Option<String>,
pub match_mode: Option<CompletionAssistantMatchMode>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompletionAssistantCandidate {
pub name: String,
pub kind: CompletionAssistantCandidateKind,
pub database: Option<String>,
pub schema: Option<String>,
pub parent_schema: Option<String>,
pub parent_name: Option<String>,
pub comment: Option<String>,
pub data_type: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompletionAssistantResponse {
pub candidates: Vec<CompletionAssistantCandidate>,
pub incomplete: bool,
pub fallback_used: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct QueryResult {
pub columns: Vec<String>,

View File

@ -0,0 +1,67 @@
use std::time::Duration;
#[tokio::test]
#[ignore = "requires DBX_LIVE_SQLSERVER_URL or DBX_LIVE_SQLSERVER_HOST/PORT/USER/PASSWORD pointing at a writable SQL Server database"]
async fn live_sqlserver_completion_assistant_searches_metadata_before_limiting() {
let database = std::env::var("DBX_LIVE_SQLSERVER_DATABASE").unwrap_or_else(|_| "tempdb".to_string());
let host = std::env::var("DBX_LIVE_SQLSERVER_HOST").unwrap_or_else(|_| "127.0.0.1".to_string());
let port = std::env::var("DBX_LIVE_SQLSERVER_PORT").ok().and_then(|value| value.parse().ok()).unwrap_or(1433);
let user = std::env::var("DBX_LIVE_SQLSERVER_USER").unwrap_or_else(|_| "sa".to_string());
let password = std::env::var("DBX_LIVE_SQLSERVER_PASSWORD").expect("DBX_LIVE_SQLSERVER_PASSWORD");
let mut client =
dbx_core::db::sqlserver::connect(&host, port, &user, &password, Some(&database), Duration::from_secs(10))
.await
.expect("connect SQL Server");
let suffix = uuid::Uuid::new_v4().simple().to_string();
let schema = format!("dbx_completion_{suffix}");
let prefix = format!("needle_{suffix}");
let table = format!("{prefix}_table");
let setup = format!(
"CREATE SCHEMA [{schema}]; CREATE TABLE [{schema}].[{table}] (id INT NOT NULL, display_name NVARCHAR(64) NULL);"
);
dbx_core::db::sqlserver::execute_query(&mut client, &setup).await.expect("create live test objects");
let request = dbx_core::types::CompletionAssistantRequest {
connection_id: "live-sqlserver".to_string(),
database: database.clone(),
schema: Some(schema.clone()),
object_kinds: vec![dbx_core::types::CompletionAssistantObjectKind::Table],
mask: prefix.clone(),
case_sensitive: false,
global_search: false,
max_results: Some(5),
search_in_comments: false,
search_in_definitions: false,
parent_schema: Some(schema.clone()),
parent_name: None,
match_mode: Some(dbx_core::types::CompletionAssistantMatchMode::Prefix),
};
let response = dbx_core::db::sqlserver::completion_assistant_search(&mut client, &request)
.await
.expect("completion assistant tables");
assert!(response
.candidates
.iter()
.any(|candidate| candidate.name == table && candidate.schema.as_deref() == Some(schema.as_str())));
let column_response = dbx_core::db::sqlserver::completion_assistant_search(
&mut client,
&dbx_core::types::CompletionAssistantRequest {
object_kinds: vec![dbx_core::types::CompletionAssistantObjectKind::Column],
mask: "display".to_string(),
parent_name: Some(table.clone()),
..request
},
)
.await
.expect("completion assistant columns");
assert!(column_response
.candidates
.iter()
.any(|candidate| candidate.name == "display_name" && candidate.parent_name.as_deref() == Some(table.as_str())));
let cleanup = format!("DROP TABLE [{schema}].[{table}]; DROP SCHEMA [{schema}];");
let _ = dbx_core::db::sqlserver::execute_query(&mut client, &cleanup).await;
}

View File

@ -204,6 +204,7 @@ async fn main() {
.route("/schema/objects", get(routes::schema::list_objects))
.route("/schema/object-statistics", get(routes::schema::list_object_statistics))
.route("/schema/completion-objects", get(routes::schema::list_completion_objects))
.route("/schema/completion-assistant", post(routes::schema::completion_assistant_search))
.route("/schema/object-source", get(routes::schema::get_object_source))
.route("/schema/columns", get(routes::schema::list_columns))
.route("/schema/indexes", get(routes::schema::list_indexes))

View File

@ -153,6 +153,14 @@ pub async fn list_completion_objects(
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
}
pub async fn completion_assistant_search(
State(state): State<Arc<WebState>>,
Json(request): Json<dbx_core::db::CompletionAssistantRequest>,
) -> Result<Json<dbx_core::db::CompletionAssistantResponse>, AppError> {
let result = dbx_core::schema::completion_assistant_search_core(&state.app, request).await.map_err(AppError)?;
Ok(Json(result))
}
pub async fn get_object_source(
State(state): State<Arc<WebState>>,
Query(q): Query<SchemaQuery>,

View File

@ -137,6 +137,14 @@ pub async fn list_completion_objects(
dbx_core::schema::list_completion_objects_core(&state, &connection_id, &database, &schema).await
}
#[tauri::command]
pub async fn completion_assistant_search(
state: State<'_, Arc<AppState>>,
request: db::CompletionAssistantRequest,
) -> Result<db::CompletionAssistantResponse, String> {
dbx_core::schema::completion_assistant_search_core(&state, request).await
}
#[tauri::command]
pub async fn get_object_source(
state: State<'_, Arc<AppState>>,

View File

@ -454,6 +454,7 @@ pub fn run() {
commands::schema::list_objects,
commands::schema::list_object_statistics,
commands::schema::list_completion_objects,
commands::schema::completion_assistant_search,
commands::schema::get_object_source,
commands::schema::list_schemas,
commands::schema::get_columns,