Skip to content

高级用法

本指南将介绍 SuperSQL 的高级特性和用法。

RAG 配置

RAG(检索增强生成)是 SuperSQL 的核心技术,通过检索相关数据来增强生成效果。

自定义 RAG 参数

java
@GetMapping("/advanced-rag")
public Object advancedRag(@RequestParam String question) {
    String sql = sqlEngine.setChatModel(chatModel)
            .setOptions(RagOptions.builder()
                    .topN(10)              // 增加检索数量
                    .rerank(true)           // 启用重排序
                    .limitScore(0.3)        // 降低分数阈值
                    .build())
            .generateSql(question);
    
    return sqlEngine.executeSql(sql);
}

RAG 工作原理

  1. 向量化: 将自然语言问题转换为向量
  2. 检索: 在向量数据库中检索最相关的表结构信息
  3. 重排序: 使用重排序模型对检索结果进行重新排序
  4. 生成: 基于检索结果生成 SQL 查询

ReRank 重排序

ReRank 可以提高检索结果的准确性,从而生成更精确的 SQL。

启用 ReRank

在配置文件中启用:

yaml
spring:
  ai:
    reranker:
      enabled: true
      model: Qwen3-Reranker-8B
      base-url: https://ai.gitee.com/v1/rerank
      api-key: your-api-key

代码中使用

java
@GetMapping("/rerank-query")
public Object rerankQuery(@RequestParam String question) {
    String sql = sqlEngine.setChatModel(chatModel)
            .setOptions(RagOptions.builder()
                    .topN(10)
                    .rerank(true)           // 启用重排序
                    .limitScore(0.3)
                    .build())
            .generateSql(question);
    
    return sqlEngine.executeSql(sql);
}

自定义 Prompt

你可以自定义生成 SQL 的 Prompt 模板。

使用 SqlpromptBuilder

java
@GetMapping("/custom-prompt")
public Object customPrompt(@RequestParam String question) {
    String customPrompt = """
        你是一个 SQL 专家。根据以下数据库表结构,生成准确的 SQL 查询。
        
        表结构:
        {table_schema}
        
        用户问题:{question}
        
        请只返回 SQL 查询语句,不要包含任何解释。
        """;
    
    String sql = sqlEngine.setChatModel(chatModel)
            .setPrompt(SqlpromptBuilder.builder()
                    .template(customPrompt)
                    .build())
            .generateSql(question);
    
    return sqlEngine.executeSql(sql);
}

Prompt 变量

变量说明
{table_schema}数据库表结构信息
{question}用户的问题
{examples}训练的 SQL 示例

性能优化

缓存 SQL 查询

java
@RestController
@RequiredArgsConstructor
public class CachedQueryController {

    private final SpringSqlEngine sqlEngine;
    private final ChatModel chatModel;
    private final CacheManager cacheManager;
    
    @GetMapping("/cached-query")
    public Object cachedQuery(@RequestParam String question) {
        Cache cache = cacheManager.getCache("sql-cache");
        String cachedSql = cache.get(question, String.class);
        
        if (cachedSql != null) {
            return sqlEngine.executeSql(cachedSql);
        }
        
        String sql = sqlEngine.setChatModel(chatModel)
                .generateSql(question);
        
        cache.put(question, sql);
        return sqlEngine.executeSql(sql);
    }
}

批量查询

java
@PostMapping("/batch-query")
public List<?> batchQuery(@RequestBody List<String> questions) {
    List<Object> results = new ArrayList<>();
    
    for (String question : questions) {
        String sql = sqlEngine.setChatModel(chatModel)
                .generateSql(question);
        Object result = sqlEngine.executeSql(sql);
        results.add(result);
    }
    
    return results;
}

异步查询

java
@GetMapping("/async-query")
public CompletableFuture<Object> asyncQuery(@RequestParam String question) {
    return CompletableFuture.supplyAsync(() -> {
        String sql = sqlEngine.setChatModel(chatModel)
                .generateSql(question);
        return sqlEngine.executeSql(sql);
    });
}

安全性

SQL 注入防护

java
@GetMapping("/safe-query")
public Object safeQuery(@RequestParam String question) {
    String sql = sqlEngine.setChatModel(chatModel)
            .generateSql(question);
    
    if (containsDangerousOperations(sql)) {
        throw new SecurityException("检测到危险操作");
    }
    
    return sqlEngine.executeSql(sql);
}

private boolean containsDangerousOperations(String sql) {
    String lowerSql = sql.toLowerCase();
    return lowerSql.contains("drop") ||
           lowerSql.contains("truncate") ||
           lowerSql.contains("delete") ||
           lowerSql.contains("alter");
}

查询权限控制

java
@GetMapping("/authorized-query")
public Object authorizedQuery(
        @RequestParam String question,
        @AuthenticationPrincipal User user) {
    
    if (!user.hasPermission("sql_query")) {
        throw new AccessDeniedException("无权限执行 SQL 查询");
    }
    
    String sql = sqlEngine.setChatModel(chatModel)
            .generateSql(question);
    
    return sqlEngine.executeSql(sql);
}

监控和日志

查询性能监控

java
@Aspect
@Component
public class QueryMonitorAspect {

    @Around("execution(* com.aispace.supersql.engine.SqlEngine.generateSql(..))")
    public Object monitorQuery(ProceedingJoinPoint joinPoint) throws Throwable {
        long startTime = System.currentTimeMillis();
        
        try {
            Object result = joinPoint.proceed();
            long duration = System.currentTimeMillis() - startTime;
            
            log.info("SQL generation completed in {} ms", duration);
            return result;
        } catch (Exception e) {
            log.error("SQL generation failed", e);
            throw e;
        }
    }
}

查询日志记录

java
@Service
@RequiredArgsConstructor
public class QueryLoggingService {

    private final SpringSqlEngine sqlEngine;
    private final QueryLogRepository logRepository;
    
    public Object queryWithLogging(String question, User user) {
        long startTime = System.currentTimeMillis();
        
        try {
            String sql = sqlEngine.setChatModel(chatModel)
                    .generateSql(question);
            Object result = sqlEngine.executeSql(sql);
            
            QueryLog log = QueryLog.builder()
                    .user(user.getId())
                    .question(question)
                    .sql(sql)
                    .duration(System.currentTimeMillis() - startTime)
                    .success(true)
                    .build();
            
            logRepository.save(log);
            return result;
            
        } catch (Exception e) {
            QueryLog log = QueryLog.builder()
                    .user(user.getId())
                    .question(question)
                    .duration(System.currentTimeMillis() - startTime)
                    .success(false)
                    .errorMessage(e.getMessage())
                    .build();
            
            logRepository.save(log);
            throw e;
        }
    }
}

多模型支持

动态切换模型

java
@RestController
@RequiredArgsConstructor
public class MultiModelController {

    private final SpringSqlEngine sqlEngine;
    private final Map<String, ChatModel> chatModels;
    
    @GetMapping("/query-with-model")
    public Object queryWithModel(
            @RequestParam String question,
            @RequestParam String modelName) {
        
        ChatModel model = chatModels.get(modelName);
        if (model == null) {
            throw new IllegalArgumentException("模型不存在: " + modelName);
        }
        
        String sql = sqlEngine.setChatModel(model)
                .generateSql(question);
        
        return sqlEngine.executeSql(sql);
    }
}

模型对比

java
@PostMapping("/compare-models")
public Map<String, Object> compareModels(
        @RequestParam String question,
        @RequestBody List<String> modelNames) {
    
    Map<String, Object> results = new HashMap<>();
    
    for (String modelName : modelNames) {
        ChatModel model = chatModels.get(modelName);
        if (model != null) {
            String sql = sqlEngine.setChatModel(model)
                    .generateSql(question);
            Object result = sqlEngine.executeSql(sql);
            results.put(modelName, Map.of(
                    "sql", sql,
                    "result", result
            ));
        }
    }
    
    return results;
}

下一步

掌握了高级用法后,你可以:

最近更新

基于 Apache 2.0 许可证发布