feat(harness): add readonly mysql tool

This commit is contained in:
zhuyongxin
2026-07-21 21:24:42 +08:00
parent 3e602781d6
commit 85029d96a7
31 changed files with 1896 additions and 53 deletions
@@ -0,0 +1,152 @@
package com.superbiz.agent.config;
import com.superbiz.agent.harness.tool.mysql.MysqlDataSourceDefinition;
import com.superbiz.agent.harness.tool.mysql.MysqlToolLimits;
import org.springframework.boot.context.properties.ConfigurationProperties;
import org.springframework.context.annotation.Configuration;
import java.util.LinkedHashMap;
import java.util.Map;
import java.util.Set;
import java.util.stream.Collectors;
/** Independent logical datasource configuration for the Agent-facing MySQL Tool. */
@Configuration
@ConfigurationProperties(prefix = "harness.mysql-tools")
public class MysqlToolProperties {
private Map<String, DataSourceProperties> dataSources = new LinkedHashMap<>();
public Map<String, DataSourceProperties> getDataSources() {
return dataSources;
}
public void setDataSources(Map<String, DataSourceProperties> dataSources) {
this.dataSources = dataSources == null ? new LinkedHashMap<>() : dataSources;
}
public Map<String, MysqlDataSourceDefinition> definitions() {
return dataSources.entrySet().stream().collect(Collectors.toUnmodifiableMap(
Map.Entry::getKey,
entry -> entry.getValue().toDefinition(entry.getKey())));
}
public static class DataSourceProperties {
private String jdbcUrl;
private String username;
private String password;
private String defaultSchema;
private int queryTimeoutSeconds = 5;
private int maxRows = 100;
private int maxCellChars = 2_000;
private int maxResultBytes = 64 * 1024;
private Map<String, SchemaProperties> allowedSchemas = new LinkedHashMap<>();
public MysqlDataSourceDefinition toDefinition(String id) {
Map<String, Map<String, Set<String>>> schemas = allowedSchemas.entrySet().stream()
.collect(Collectors.toMap(
Map.Entry::getKey,
entry -> entry.getValue().tables.entrySet().stream()
.collect(Collectors.toMap(Map.Entry::getKey,
table -> Set.copyOf(table.getValue().columns)))));
return new MysqlDataSourceDefinition(id, defaultSchema, schemas,
new MysqlToolLimits(maxRows, maxCellChars, maxResultBytes, queryTimeoutSeconds));
}
public String getJdbcUrl() {
return jdbcUrl;
}
public void setJdbcUrl(String jdbcUrl) {
this.jdbcUrl = jdbcUrl;
}
public String getUsername() {
return username;
}
public void setUsername(String username) {
this.username = username;
}
public String getPassword() {
return password;
}
public void setPassword(String password) {
this.password = password;
}
public String getDefaultSchema() {
return defaultSchema;
}
public void setDefaultSchema(String defaultSchema) {
this.defaultSchema = defaultSchema;
}
public int getQueryTimeoutSeconds() {
return queryTimeoutSeconds;
}
public void setQueryTimeoutSeconds(int queryTimeoutSeconds) {
this.queryTimeoutSeconds = queryTimeoutSeconds;
}
public int getMaxRows() {
return maxRows;
}
public void setMaxRows(int maxRows) {
this.maxRows = maxRows;
}
public int getMaxCellChars() {
return maxCellChars;
}
public void setMaxCellChars(int maxCellChars) {
this.maxCellChars = maxCellChars;
}
public int getMaxResultBytes() {
return maxResultBytes;
}
public void setMaxResultBytes(int maxResultBytes) {
this.maxResultBytes = maxResultBytes;
}
public Map<String, SchemaProperties> getAllowedSchemas() {
return allowedSchemas;
}
public void setAllowedSchemas(Map<String, SchemaProperties> allowedSchemas) {
this.allowedSchemas = allowedSchemas == null ? new LinkedHashMap<>() : allowedSchemas;
}
}
public static class SchemaProperties {
private Map<String, TableProperties> tables = new LinkedHashMap<>();
public Map<String, TableProperties> getTables() {
return tables;
}
public void setTables(Map<String, TableProperties> tables) {
this.tables = tables == null ? new LinkedHashMap<>() : tables;
}
}
public static class TableProperties {
private Set<String> columns = Set.of();
public Set<String> getColumns() {
return columns;
}
public void setColumns(Set<String> columns) {
this.columns = columns == null ? Set.of() : columns;
}
}
}
@@ -0,0 +1,52 @@
package com.superbiz.agent.harness.tool.adapter;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.tool.boundary.ToolBoundary;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryErrorCode;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryResult;
import com.superbiz.agent.harness.tool.boundary.ToolCallRequestEnvelope;
import com.superbiz.agent.harness.tool.contract.MysqlToolRequest;
import com.superbiz.agent.harness.tool.mysql.MysqlQueryPlan;
import com.superbiz.agent.harness.tool.mysql.MysqlReadOnlyExecutor;
import com.superbiz.agent.harness.tool.mysql.MysqlResultProjector;
import com.superbiz.agent.harness.tool.mysql.MysqlSecurityException;
import com.superbiz.agent.harness.tool.mysql.MysqlSqlValidator;
import java.util.Objects;
/** Validates and runs the logical MySQL Tool through the canonical boundary. */
public final class MysqlToolAdapter {
private final ToolBoundary boundary;
private final ObjectMapper objectMapper;
private final MysqlSqlValidator validator;
private final MysqlReadOnlyExecutor executor;
private final MysqlResultProjector projector;
public MysqlToolAdapter(ToolBoundary boundary, ObjectMapper objectMapper,
MysqlSqlValidator validator, MysqlReadOnlyExecutor executor,
MysqlResultProjector projector) {
this.boundary = Objects.requireNonNull(boundary, "boundary must not be null");
this.objectMapper = Objects.requireNonNull(objectMapper, "objectMapper must not be null");
this.validator = Objects.requireNonNull(validator, "validator must not be null");
this.executor = Objects.requireNonNull(executor, "executor must not be null");
this.projector = Objects.requireNonNull(projector, "projector must not be null");
}
public ToolBoundaryResult execute(RunContext context, ToolCallRequestEnvelope envelope) {
try {
MysqlToolRequest request = objectMapper.readValue(envelope.requestJson(), MysqlToolRequest.class);
MysqlQueryPlan plan = validator.validate(request);
return boundary.execute(context, envelope,
ignored -> objectMapper.writeValueAsString(executor.execute(plan, context)),
raw -> projector.project(request, envelope.toolCallId(), raw, plan.dataSource().limits()));
} catch (MysqlSecurityException | IllegalArgumentException e) {
return ToolBoundaryResult.error(envelope == null ? null : envelope.toolCallId(),
ToolBoundaryErrorCode.INVALID_REQUEST);
} catch (Exception e) {
return ToolBoundaryResult.error(envelope == null ? null : envelope.toolCallId(),
ToolBoundaryErrorCode.INVALID_REQUEST);
}
}
}
@@ -0,0 +1,140 @@
package com.superbiz.agent.harness.tool.mysql;
import com.superbiz.agent.harness.core.RunContext;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import javax.sql.DataSource;
import java.nio.charset.StandardCharsets;
import java.sql.Connection;
import java.sql.PreparedStatement;
import java.sql.ResultSet;
import java.sql.ResultSetMetaData;
import java.sql.SQLException;
import java.sql.Statement;
import java.time.Clock;
import java.util.Base64;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.concurrent.atomic.AtomicReference;
/** JDBC implementation with read-only, timeout, row and cancellation controls. */
public final class JdbcMysqlReadOnlyExecutor implements MysqlReadOnlyExecutor {
private static final Logger log = LoggerFactory.getLogger(JdbcMysqlReadOnlyExecutor.class);
private final Map<String, DataSource> dataSources;
private final Clock clock;
public JdbcMysqlReadOnlyExecutor(Map<String, DataSource> dataSources, Clock clock) {
this.dataSources = Map.copyOf(dataSources);
this.clock = Objects.requireNonNull(clock, "clock must not be null");
}
@Override
public MysqlRawResult execute(MysqlQueryPlan plan, RunContext context) throws Exception {
DataSource dataSource = dataSources.get(plan.dataSource().id());
if (dataSource == null) {
throw new MysqlSecurityException("logical data source is not configured");
}
MysqlToolLimits limits = plan.dataSource().limits();
try (Connection connection = dataSource.getConnection()) {
connection.setReadOnly(true);
try (PreparedStatement statement = connection.prepareStatement(
plan.normalizedSql(), ResultSet.TYPE_FORWARD_ONLY, ResultSet.CONCUR_READ_ONLY)) {
statement.setQueryTimeout(limits.queryTimeoutSeconds());
statement.setMaxRows(limits.maxRows() + 1);
bind(statement, plan.params());
AtomicReference<Statement> statementRef = new AtomicReference<>(statement);
context.cancellation().onCancel(ignored -> cancel(statementRef.get()));
checkRun(context);
try (ResultSet resultSet = statement.executeQuery()) {
ResultSetMetaData metadata = resultSet.getMetaData();
java.util.ArrayList<String> columns = new java.util.ArrayList<>();
for (int i = 1; i <= metadata.getColumnCount(); i++) {
columns.add(metadata.getColumnLabel(i));
}
java.util.ArrayList<Map<String, Object>> rows = new java.util.ArrayList<>();
boolean truncated = false;
while (resultSet.next()) {
checkRun(context);
if (rows.size() >= limits.maxRows()) {
truncated = true;
break;
}
Map<String, Object> row = new LinkedHashMap<>();
for (int i = 1; i <= metadata.getColumnCount(); i++) {
String column = metadata.getColumnLabel(i);
CellValue cell = jsonSafe(resultSet.getObject(i), limits.maxCellChars());
row.put(column, cell.value());
truncated |= cell.truncated();
}
rows.add(row);
if (estimatedBytes(rows) > limits.maxResultBytes()) {
rows.remove(rows.size() - 1);
truncated = true;
break;
}
}
return new MysqlRawResult(columns, rows, truncated);
} finally {
statementRef.set(null);
}
}
} catch (MysqlSecurityException e) {
throw e;
} catch (SQLException e) {
log.debug("MySQL read-only execution failed: sqlState={}", e.getSQLState());
throw new SQLException("read-only query failed", e);
}
}
private static void bind(PreparedStatement statement, List<Object> params) throws SQLException {
for (int i = 0; i < params.size(); i++) {
statement.setObject(i + 1, params.get(i));
}
}
private void checkRun(RunContext context) throws SQLException {
if (context.cancellation().isCancelled() || !clock.instant().isBefore(context.deadline())) {
throw new SQLException("run cancelled or deadline exceeded");
}
}
private static void cancel(Statement statement) {
if (statement == null) {
return;
}
try {
statement.cancel();
} catch (SQLException e) {
log.debug("Unable to cancel MySQL statement", e);
}
}
private static CellValue jsonSafe(Object value, int maxCellChars) {
if (value == null || value instanceof Number || value instanceof Boolean) {
return new CellValue(value, false);
}
String text;
if (value instanceof byte[] bytes) {
text = Base64.getEncoder().encodeToString(bytes);
} else {
text = String.valueOf(value);
}
return text.length() <= maxCellChars
? new CellValue(text, false)
: new CellValue(text.substring(0, maxCellChars), true);
}
private static int estimatedBytes(List<Map<String, Object>> rows) {
return rows.toString().getBytes(StandardCharsets.UTF_8).length;
}
private record CellValue(Object value, boolean truncated) {
}
}
@@ -0,0 +1,51 @@
package com.superbiz.agent.harness.tool.mysql;
import java.util.Collections;
import java.util.LinkedHashMap;
import java.util.Map;
import java.util.Objects;
import java.util.Set;
import java.util.TreeSet;
/** Logical datasource metadata and exact schema/table/column authorization. */
public record MysqlDataSourceDefinition(
String id,
String defaultSchema,
Map<String, Map<String, Set<String>>> allowedSchemas,
MysqlToolLimits limits) {
public MysqlDataSourceDefinition {
requireText(id, "id");
requireText(defaultSchema, "defaultSchema");
Objects.requireNonNull(allowedSchemas, "allowedSchemas must not be null");
Objects.requireNonNull(limits, "limits must not be null");
Map<String, Map<String, Set<String>>> schemas = new LinkedHashMap<>();
allowedSchemas.forEach((schema, tables) -> {
requireText(schema, "schema");
Map<String, Set<String>> copiedTables = new LinkedHashMap<>();
tables.forEach((table, columns) -> {
requireText(table, "table");
copiedTables.put(table, Collections.unmodifiableSet(new TreeSet<>(columns)));
});
schemas.put(schema, Collections.unmodifiableMap(copiedTables));
});
allowedSchemas = Collections.unmodifiableMap(schemas);
if (!allowedSchemas.containsKey(defaultSchema)) {
throw new IllegalArgumentException("defaultSchema must be allowlisted");
}
}
public boolean allowsTable(String schema, String table) {
return allowedSchemas.containsKey(schema) && allowedSchemas.get(schema).containsKey(table);
}
public boolean allowsColumn(String schema, String table, String column) {
return allowsTable(schema, table) && allowedSchemas.get(schema).get(table).contains(column);
}
private static void requireText(String value, String name) {
if (value == null || value.isBlank()) {
throw new IllegalArgumentException(name + " must not be blank");
}
}
}
@@ -0,0 +1,16 @@
package com.superbiz.agent.harness.tool.mysql;
import com.superbiz.agent.harness.tool.contract.MysqlToolRequest;
import java.util.List;
public record MysqlQueryPlan(
MysqlToolRequest request,
MysqlDataSourceDefinition dataSource,
String normalizedSql,
List<Object> params) {
public MysqlQueryPlan {
params = List.copyOf(params);
}
}
@@ -0,0 +1,20 @@
package com.superbiz.agent.harness.tool.mysql;
import java.util.List;
import java.util.Collections;
import java.util.LinkedHashMap;
import java.util.Map;
/** Harness-only raw query result; never returned directly to an Agent. */
public record MysqlRawResult(
List<String> columns,
List<Map<String, Object>> rows,
boolean truncated) {
public MysqlRawResult {
columns = List.copyOf(columns);
rows = rows.stream()
.map(row -> Collections.unmodifiableMap(new LinkedHashMap<>(row)))
.toList();
}
}
@@ -0,0 +1,8 @@
package com.superbiz.agent.harness.tool.mysql;
import com.superbiz.agent.harness.core.RunContext;
@FunctionalInterface
public interface MysqlReadOnlyExecutor {
MysqlRawResult execute(MysqlQueryPlan plan, RunContext context) throws Exception;
}
@@ -0,0 +1,133 @@
package com.superbiz.agent.harness.tool.mysql;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.contract.EvidenceStatus;
import com.superbiz.agent.harness.tool.boundary.ProjectedToolResult;
import com.superbiz.agent.harness.tool.contract.MysqlToolRequest;
import com.superbiz.agent.harness.tool.contract.MysqlToolResult;
import java.nio.charset.StandardCharsets;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Locale;
import java.util.Map;
/** Projects raw JDBC rows into the bounded Agent-facing MySQL contract. */
public final class MysqlResultProjector {
private static final List<String> SENSITIVE_TOKENS = List.of(
"password", "passwd", "token", "secret", "api_key", "apikey", "credential");
private final ObjectMapper objectMapper;
private final MysqlToolLimits limits;
public MysqlResultProjector(ObjectMapper objectMapper) {
this(objectMapper, MysqlToolLimits.defaults());
}
public MysqlResultProjector(ObjectMapper objectMapper, MysqlToolLimits limits) {
this.objectMapper = objectMapper;
this.limits = limits;
}
public ProjectedToolResult project(MysqlToolRequest request, String toolCallId,
String rawResponse) throws Exception {
return project(request, toolCallId, rawResponse, limits);
}
public ProjectedToolResult project(MysqlToolRequest request, String toolCallId,
String rawResponse, MysqlToolLimits projectionLimits) throws Exception {
if (request == null || toolCallId == null || toolCallId.isBlank()) {
throw new IllegalArgumentException("request and tool call ID are required");
}
JsonNode root = objectMapper.readTree(rawResponse);
if (root == null || !root.isObject() || !root.path("columns").isArray()
|| !root.path("rows").isArray()) {
throw new IllegalArgumentException("MySQL raw result is invalid");
}
List<String> columns = new ArrayList<>();
java.util.LinkedHashSet<String> uniqueColumns = new java.util.LinkedHashSet<>();
root.path("columns").forEach(node -> {
String column = node.asText();
if (column.isBlank() || !uniqueColumns.add(column)) {
throw new IllegalArgumentException("MySQL columns must be unique and non-blank");
}
columns.add(column);
});
List<Map<String, Object>> rows = new ArrayList<>();
boolean truncated = root.path("truncated").asBoolean(false);
for (JsonNode rowNode : root.path("rows")) {
if (rows.size() >= projectionLimits.maxRows()) {
truncated = true;
break;
}
Map<String, Object> row = new LinkedHashMap<>();
for (String column : columns) {
JsonNode value = rowNode.get(column);
CellProjection cell = projectCell(column, value, projectionLimits.maxCellChars());
row.put(column, cell.value());
truncated |= cell.truncated();
}
rows.add(row);
if (utf8Bytes(rows.toString()) > projectionLimits.maxResultBytes()) {
rows.remove(rows.size() - 1);
truncated = true;
break;
}
}
MysqlToolResult result = new MysqlToolResult(
rows.isEmpty() ? EvidenceStatus.NO_EVIDENCE : EvidenceStatus.EVIDENCE_FOUND,
toolCallId, columns, rows, rows.size(), truncated);
result = fitBudget(result, projectionLimits.maxResultBytes());
return new ProjectedToolResult(objectMapper.writeValueAsString(result), result.evidenceStatus());
}
private MysqlToolResult fitBudget(MysqlToolResult result, int maxResultBytes) throws Exception {
MysqlToolResult current = result;
while (utf8Bytes(objectMapper.writeValueAsString(current)) > maxResultBytes
&& !current.rows().isEmpty()) {
List<Map<String, Object>> reduced = new ArrayList<>(current.rows());
reduced.remove(reduced.size() - 1);
current = new MysqlToolResult(
reduced.isEmpty() ? EvidenceStatus.NO_EVIDENCE : EvidenceStatus.EVIDENCE_FOUND,
current.toolCallId(), current.columns(), reduced, reduced.size(), true);
}
if (utf8Bytes(objectMapper.writeValueAsString(current)) > maxResultBytes) {
throw new IllegalArgumentException("MySQL projection exceeds total budget");
}
return current;
}
private CellProjection projectCell(String column, JsonNode value, int maxCellChars) {
if (value == null || value.isNull()) {
return new CellProjection(null, false);
}
if (isSensitive(column)) {
return new CellProjection("[REDACTED]", true);
}
if (value.isNumber()) {
return new CellProjection(value.numberValue(), false);
}
if (value.isBoolean()) {
return new CellProjection(value.booleanValue(), false);
}
String text = value.isTextual() ? value.textValue() : value.toString();
return text.length() <= maxCellChars
? new CellProjection(text, false)
: new CellProjection(text.substring(0, maxCellChars), true);
}
private static boolean isSensitive(String column) {
String normalized = column == null ? "" : column.toLowerCase(Locale.ROOT);
return SENSITIVE_TOKENS.stream().anyMatch(normalized::contains);
}
private static int utf8Bytes(String value) {
return value.getBytes(StandardCharsets.UTF_8).length;
}
private record CellProjection(Object value, boolean truncated) {
}
}
@@ -0,0 +1,12 @@
package com.superbiz.agent.harness.tool.mysql;
public final class MysqlSecurityException extends RuntimeException {
public MysqlSecurityException(String message) {
super(message);
}
public MysqlSecurityException(String message, Throwable cause) {
super(message, cause);
}
}
@@ -0,0 +1,356 @@
package com.superbiz.agent.harness.tool.mysql;
import com.superbiz.agent.harness.tool.contract.MysqlToolRequest;
import net.sf.jsqlparser.expression.AnalyticExpression;
import net.sf.jsqlparser.expression.CaseExpression;
import net.sf.jsqlparser.expression.Function;
import net.sf.jsqlparser.expression.DateValue;
import net.sf.jsqlparser.expression.DoubleValue;
import net.sf.jsqlparser.expression.HexValue;
import net.sf.jsqlparser.expression.JdbcParameter;
import net.sf.jsqlparser.expression.LongValue;
import net.sf.jsqlparser.expression.OracleHierarchicalExpression;
import net.sf.jsqlparser.expression.StringValue;
import net.sf.jsqlparser.expression.TimeValue;
import net.sf.jsqlparser.expression.TimestampValue;
import net.sf.jsqlparser.expression.operators.relational.ExistsExpression;
import net.sf.jsqlparser.schema.Column;
import net.sf.jsqlparser.parser.CCJSqlParserUtil;
import net.sf.jsqlparser.schema.Table;
import net.sf.jsqlparser.statement.Statement;
import net.sf.jsqlparser.statement.Statements;
import net.sf.jsqlparser.statement.select.AllColumns;
import net.sf.jsqlparser.statement.select.AllTableColumns;
import net.sf.jsqlparser.statement.select.FromItem;
import net.sf.jsqlparser.statement.select.Join;
import net.sf.jsqlparser.statement.select.PlainSelect;
import net.sf.jsqlparser.statement.select.Select;
import net.sf.jsqlparser.statement.select.SelectBody;
import net.sf.jsqlparser.statement.select.SelectExpressionItem;
import net.sf.jsqlparser.statement.select.SetOperationList;
import net.sf.jsqlparser.statement.select.SubSelect;
import net.sf.jsqlparser.statement.values.ValuesStatement;
import net.sf.jsqlparser.expression.Expression;
import net.sf.jsqlparser.expression.ExpressionVisitorAdapter;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Locale;
import java.util.Map;
import java.util.Objects;
import java.util.Set;
/** Fail-closed SQL policy for the Agent-facing MySQL Tool. */
public final class MysqlSqlValidator {
private static final Set<String> ALLOWED_FUNCTIONS = Set.of("COUNT", "SUM", "AVG", "MIN", "MAX");
private final Map<String, MysqlDataSourceDefinition> dataSources;
public MysqlSqlValidator(Map<String, MysqlDataSourceDefinition> dataSources) {
Objects.requireNonNull(dataSources, "dataSources must not be null");
this.dataSources = Map.copyOf(dataSources);
}
public MysqlQueryPlan validate(MysqlToolRequest request) {
if (request == null || request.dataSource() == null || request.dataSource().isBlank()
|| request.sql() == null || request.sql().isBlank()) {
throw new MysqlSecurityException("data_source and sql are required");
}
MysqlDataSourceDefinition dataSource = dataSources.get(request.dataSource());
if (dataSource == null) {
throw new MysqlSecurityException("unknown logical data source");
}
if (request.sql().length() > 16_384) {
throw new MysqlSecurityException("SQL exceeds policy length");
}
try {
Statements statements = CCJSqlParserUtil.parseStatements(request.sql());
if (statements.getStatements() == null || statements.getStatements().size() != 1) {
throw new MysqlSecurityException("exactly one SQL statement is required");
}
Statement statement = statements.getStatements().get(0);
if (!(statement instanceof Select select)) {
throw new MysqlSecurityException("only SELECT is allowed");
}
if (select.getWithItemsList() != null && !select.getWithItemsList().isEmpty()) {
throw new MysqlSecurityException("WITH is not allowed");
}
SelectBody body = select.getSelectBody();
if (!(body instanceof PlainSelect plainSelect)
|| body instanceof SetOperationList
|| body instanceof ValuesStatement) {
throw new MysqlSecurityException("only a plain SELECT is allowed");
}
if (plainSelect.isForUpdate() || plainSelect.isSkipLocked()
|| plainSelect.getFromItem() == null) {
throw new MysqlSecurityException("locking or missing FROM is not allowed");
}
if (plainSelect.getIntoTables() != null && !plainSelect.getIntoTables().isEmpty()
|| plainSelect.getOffset() != null || plainSelect.getFetch() != null
|| plainSelect.getTop() != null || plainSelect.getFirst() != null
|| plainSelect.getSkip() != null || plainSelect.getOptimizeFor() != null
|| plainSelect.getOracleHierarchical() != null
|| plainSelect.getKsqlWindow() != null
|| plainSelect.getWindowDefinitions() != null && !plainSelect.getWindowDefinitions().isEmpty()) {
throw new MysqlSecurityException("unsupported SELECT clause");
}
Map<String, TableRef> tables = new LinkedHashMap<>();
registerTable(plainSelect.getFromItem(), dataSource, tables);
List<Join> joins = plainSelect.getJoins() == null ? List.of() : plainSelect.getJoins();
for (Join join : joins) {
if (join.isCross() || join.isRight() || join.isFull() || join.isOuter()
|| (!join.isInner() && !join.isLeft())) {
throw new MysqlSecurityException("only INNER/LEFT JOIN is allowed");
}
registerTable(join.getRightItem(), dataSource, tables);
if (join.getOnExpressions() != null) {
join.getOnExpressions().forEach(expression ->
validateExpression(expression, tables, dataSource));
}
if (join.getUsingColumns() != null) {
for (Column column : join.getUsingColumns()) {
validateColumn(column, tables, dataSource);
}
}
}
if (plainSelect.getSelectItems() == null || plainSelect.getSelectItems().isEmpty()) {
throw new MysqlSecurityException("projection must be explicit");
}
for (var item : plainSelect.getSelectItems()) {
if (item instanceof AllColumns || item instanceof AllTableColumns) {
throw new MysqlSecurityException("wildcard projection is not allowed");
}
if (!(item instanceof SelectExpressionItem expressionItem)) {
throw new MysqlSecurityException("unsupported select item");
}
validateExpression(expressionItem.getExpression(), tables, dataSource);
}
validateExpression(plainSelect.getWhere(), tables, dataSource);
validateExpression(plainSelect.getHaving(), tables, dataSource);
if (plainSelect.getGroupBy() != null) {
for (Expression expression : plainSelect.getGroupBy().getGroupByExpressions()) {
validateExpression(expression, tables, dataSource);
}
}
if (plainSelect.getOrderByElements() != null) {
plainSelect.getOrderByElements().forEach(order ->
validateExpression(order.getExpression(), tables, dataSource));
}
int placeholders = countPlaceholders(plainSelect);
int provided = request.params() == null ? 0 : request.params().size();
if (placeholders != provided) {
throw new MysqlSecurityException("placeholder count does not match params");
}
return new MysqlQueryPlan(request, dataSource, statement.toString(),
request.params() == null ? List.of() : request.params());
} catch (MysqlSecurityException e) {
throw e;
} catch (Exception e) {
throw new MysqlSecurityException("SQL cannot be safely validated", e);
}
}
private void registerTable(FromItem item, MysqlDataSourceDefinition dataSource,
Map<String, TableRef> tables) {
if (!(item instanceof Table table)) {
throw new MysqlSecurityException("subqueries and non-table sources are not allowed");
}
String schema = table.getSchemaName();
if (schema == null || schema.isBlank()) {
schema = dataSource.defaultSchema();
}
String name = table.getName();
if (!dataSource.allowsTable(schema, name)) {
throw new MysqlSecurityException("table is not allowlisted");
}
String alias = table.getAlias() == null ? name : table.getAlias().getName();
String key = alias.toLowerCase(Locale.ROOT);
if (tables.putIfAbsent(key, new TableRef(schema, name)) != null) {
throw new MysqlSecurityException("duplicate table alias");
}
}
private void validateExpression(Expression expression, Map<String, TableRef> tables,
MysqlDataSourceDefinition dataSource) {
if (expression == null) {
return;
}
expression.accept(new ExpressionVisitorAdapter() {
@Override
public void visit(Column column) {
validateColumn(column, tables, dataSource);
}
@Override
public void visit(Function function) {
String name = function.getName() == null ? "" : function.getName().toUpperCase(Locale.ROOT);
if (!ALLOWED_FUNCTIONS.contains(name)) {
throw new MysqlSecurityException("function is not allowlisted");
}
boolean countStar = function.isAllColumns()
|| (function.getParameters() != null
&& function.getParameters().getExpressions() != null
&& function.getParameters().getExpressions().size() == 1
&& function.getParameters().getExpressions().get(0) instanceof AllColumns);
if (countStar && !"COUNT".equals(name)) {
throw new MysqlSecurityException("only COUNT(*) is allowed");
}
if (countStar) {
return;
}
super.visit(function);
}
@Override
public void visit(JdbcParameter parameter) {
// Counted separately by the parser walk below.
}
@Override
public void visit(StringValue value) {
throw new MysqlSecurityException("literal values must use parameters");
}
@Override
public void visit(LongValue value) {
throw new MysqlSecurityException("literal values must use parameters");
}
@Override
public void visit(DoubleValue value) {
throw new MysqlSecurityException("literal values must use parameters");
}
@Override
public void visit(HexValue value) {
throw new MysqlSecurityException("literal values must use parameters");
}
@Override
public void visit(DateValue value) {
throw new MysqlSecurityException("literal values must use parameters");
}
@Override
public void visit(TimeValue value) {
throw new MysqlSecurityException("literal values must use parameters");
}
@Override
public void visit(TimestampValue value) {
throw new MysqlSecurityException("literal values must use parameters");
}
@Override
public void visit(SubSelect subSelect) {
throw new MysqlSecurityException("subqueries are not allowed");
}
@Override
public void visit(AllColumns allColumns) {
throw new MysqlSecurityException("wildcard projection is not allowed");
}
@Override
public void visit(AllTableColumns allTableColumns) {
throw new MysqlSecurityException("wildcard projection is not allowed");
}
@Override
public void visit(AnalyticExpression analyticExpression) {
throw new MysqlSecurityException("window functions are not allowed");
}
@Override
public void visit(CaseExpression caseExpression) {
throw new MysqlSecurityException("CASE expressions are not allowed");
}
@Override
public void visit(ExistsExpression existsExpression) {
throw new MysqlSecurityException("EXISTS is not allowed");
}
@Override
public void visit(OracleHierarchicalExpression expression) {
throw new MysqlSecurityException("hierarchical expressions are not allowed");
}
});
}
private void validateColumn(Column column, Map<String, TableRef> tables,
MysqlDataSourceDefinition dataSource) {
String name = column.getColumnName();
if (name == null || name.isBlank() || "*".equals(name)) {
throw new MysqlSecurityException("invalid column");
}
Table table = column.getTable();
if (table != null && table.getName() != null && !table.getName().isBlank()) {
TableRef ref = tables.get(table.getName().toLowerCase(Locale.ROOT));
if (ref == null || !dataSource.allowsColumn(ref.schema(), ref.table(), name)) {
throw new MysqlSecurityException("column is not allowlisted");
}
return;
}
List<TableRef> matches = tables.values().stream()
.filter(ref -> dataSource.allowsColumn(ref.schema(), ref.table(), name))
.toList();
if (matches.size() != 1) {
throw new MysqlSecurityException("unqualified column is ambiguous or not allowlisted");
}
}
private static int countPlaceholders(PlainSelect plainSelect) {
// Parser assigns JdbcParameter nodes; use the canonical SQL token count only after
// the AST has been accepted, so quoted question marks are not counted.
PlaceholderCounter counter = new PlaceholderCounter();
List<Expression> expressions = new ArrayList<>();
for (var item : plainSelect.getSelectItems()) {
if (item instanceof SelectExpressionItem expressionItem) {
expressions.add(expressionItem.getExpression());
}
}
expressions.add(plainSelect.getWhere());
expressions.add(plainSelect.getHaving());
if (plainSelect.getGroupBy() != null) {
expressions.addAll(plainSelect.getGroupBy().getGroupByExpressions());
}
if (plainSelect.getOrderByElements() != null) {
plainSelect.getOrderByElements().forEach(order -> expressions.add(order.getExpression()));
}
if (plainSelect.getJoins() != null) {
plainSelect.getJoins().forEach(join -> {
if (join.getOnExpressions() != null) {
expressions.addAll(join.getOnExpressions());
}
});
}
for (Expression expression : expressions) {
if (expression != null) {
expression.accept(counter);
}
}
return counter.count;
}
private static final class PlaceholderCounter extends ExpressionVisitorAdapter {
private int count;
@Override
public void visit(JdbcParameter parameter) {
count++;
}
@Override
public void visit(SubSelect subSelect) {
throw new MysqlSecurityException("subqueries are not allowed");
}
}
private record TableRef(String schema, String table) {
}
}
@@ -0,0 +1,18 @@
package com.superbiz.agent.harness.tool.mysql;
public record MysqlToolLimits(
int maxRows,
int maxCellChars,
int maxResultBytes,
int queryTimeoutSeconds) {
public MysqlToolLimits {
if (maxRows <= 0 || maxCellChars <= 0 || maxResultBytes <= 0 || queryTimeoutSeconds <= 0) {
throw new IllegalArgumentException("MySQL limits must be positive");
}
}
public static MysqlToolLimits defaults() {
return new MysqlToolLimits(100, 2_000, 64 * 1024, 5);
}
}
+6
View File
@@ -200,3 +200,9 @@ logging:
max-file-size: 10MB # 单个日志文件最大 10MB
max-history: 30 # 保留 30 天
total-size-cap: 1GB # 所有日志文件总大小上限 1GB
# Agent-facing MySQL Tool uses independent logical datasources only.
# Production entries are supplied by a dedicated profile and Secret injection.
harness:
mysql-tools:
data-sources: {}
@@ -0,0 +1,120 @@
package com.superbiz.agent.harness.tool.adapter;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.contract.EvidenceStatus;
import com.superbiz.agent.harness.core.HarnessCoreFixtures;
import com.superbiz.agent.harness.core.MutableClock;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.tool.boundary.ToolBoundary;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryErrorCode;
import com.superbiz.agent.harness.tool.boundary.ToolBoundaryResult;
import com.superbiz.agent.harness.tool.boundary.ToolCallRequestEnvelope;
import com.superbiz.agent.harness.tool.mysql.MysqlDataSourceDefinition;
import com.superbiz.agent.harness.tool.mysql.MysqlRawResult;
import com.superbiz.agent.harness.tool.mysql.MysqlResultProjector;
import com.superbiz.agent.harness.tool.mysql.MysqlSqlValidator;
import com.superbiz.agent.harness.tool.mysql.MysqlToolLimits;
import com.superbiz.agent.harness.tool.store.CanonicalInvocationLimits;
import com.superbiz.agent.harness.tool.store.CanonicalInvocationStore;
import com.superbiz.agent.harness.tool.store.CanonicalToolInvocation;
import com.superbiz.agent.harness.tool.store.DuplicateInvocationException;
import com.superbiz.agent.harness.tool.store.ToolCallKeyFactory;
import org.junit.jupiter.api.Test;
import java.time.Duration;
import java.time.Instant;
import java.util.HashMap;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.Set;
import java.util.concurrent.atomic.AtomicInteger;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertTrue;
class MysqlToolAdapterTest {
private final ObjectMapper objectMapper = new ObjectMapper();
@Test
void reusesBoundaryAndRejectsInvalidSqlBeforeExecution() {
MutableClock clock = new MutableClock(Instant.parse("2026-07-21T10:00:00Z"));
FakeStore store = new FakeStore();
ToolBoundary boundary = new ToolBoundary(HarnessCoreFixtures.core(clock),
new ToolCallKeyFactory("superbiz:harness:tool-call"), store, objectMapper, clock);
MysqlDataSourceDefinition definition = new MysqlDataSourceDefinition(
"order_readonly", "order_db",
Map.of("order_db", Map.of("biz_order", Set.of("order_id", "password"))),
MysqlToolLimits.defaults());
AtomicInteger executions = new AtomicInteger();
MysqlToolAdapter adapter = new MysqlToolAdapter(boundary, objectMapper,
new MysqlSqlValidator(Map.of("order_readonly", definition)),
(plan, context) -> {
executions.incrementAndGet();
Map<String, Object> row = new LinkedHashMap<>();
row.put("order_id", "order-1");
row.put("password", "raw-secret");
return new MysqlRawResult(List.of("order_id", "password"), List.of(row), false);
}, new MysqlResultProjector(objectMapper));
RunContext context = HarnessCoreFixtures.core(clock).startRun("session", "run-mysql");
ToolBoundaryResult valid = adapter.execute(context, envelope("run-mysql", "framework-mysql-1",
"{\"data_source\":\"order_readonly\",\"sql\":\"SELECT order_id, password FROM order_db.biz_order WHERE order_id = ?\",\"params\":[\"order-1\"]}"));
ToolBoundaryResult invalid = adapter.execute(context, envelope("run-mysql", "framework-mysql-2",
"{\"data_source\":\"order_readonly\",\"sql\":\"SELECT * FROM order_db.biz_order\",\"params\":[]}"));
assertEquals(EvidenceStatus.EVIDENCE_FOUND, valid.evidenceStatus());
assertEquals("framework-mysql-1", valid.toolCallId());
assertFalse(valid.agentResult().contains("raw-secret"));
assertTrue(store.records.values().iterator().next().rawResponse().contains("raw-secret"));
assertEquals(ToolBoundaryErrorCode.INVALID_REQUEST.name(), invalid.errorCode());
assertEquals(1, executions.get());
}
private ToolCallRequestEnvelope envelope(String runId, String callId, String requestJson) {
return new ToolCallRequestEnvelope(runId, callId, "query_mysql", requestJson, true, true);
}
private static final class FakeStore implements CanonicalInvocationStore {
private final CanonicalInvocationLimits limits =
new CanonicalInvocationLimits(Duration.ofHours(1), 100_000, 64_000);
private final Map<String, CanonicalToolInvocation> records = new HashMap<>();
@Override
public CanonicalInvocationLimits limits() {
return limits;
}
@Override
public void begin(String key, CanonicalToolInvocation invocation) {
if (records.putIfAbsent(key, invocation) != null) {
throw new DuplicateInvocationException();
}
}
@Override
public Optional<CanonicalToolInvocation> find(String key) {
return Optional.ofNullable(records.get(key));
}
@Override
public CanonicalToolInvocation markReady(String key, String rawResponse, String agentResult,
EvidenceStatus evidenceStatus, Instant completedAt) {
CanonicalToolInvocation updated = records.get(key)
.markReady(rawResponse, agentResult, evidenceStatus, completedAt);
records.put(key, updated);
return updated;
}
@Override
public CanonicalToolInvocation markError(String key, String rawResponse,
String errorCode, Instant completedAt) {
CanonicalToolInvocation updated = records.get(key).markError(rawResponse, errorCode, completedAt);
records.put(key, updated);
return updated;
}
}
}
@@ -0,0 +1,112 @@
package com.superbiz.agent.harness.tool.mysql;
import com.superbiz.agent.harness.core.HarnessCoreFixtures;
import com.superbiz.agent.harness.core.MutableClock;
import com.superbiz.agent.harness.core.RunCancellationReason;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.tool.contract.MysqlToolRequest;
import org.h2.jdbcx.JdbcDataSource;
import org.junit.jupiter.api.Test;
import java.sql.SQLException;
import java.sql.Connection;
import java.sql.PreparedStatement;
import java.sql.ResultSet;
import java.sql.ResultSetMetaData;
import java.time.Instant;
import java.util.List;
import java.util.Map;
import java.util.Set;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertThrows;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.ArgumentMatchers.anyString;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
class JdbcMysqlReadOnlyExecutorTest {
@Test
void executesBoundPreparedSelectAndTruncatesRows() throws Exception {
JdbcDataSource dataSource = dataSource();
MysqlDataSourceDefinition definition = definition(new MysqlToolLimits(1, 20, 1024, 2));
MysqlSqlValidator validator = new MysqlSqlValidator(Map.of("order_readonly", definition));
MysqlQueryPlan plan = validator.validate(new MysqlToolRequest(
"order_readonly", "SELECT order_id, status FROM biz_order WHERE status = ? ORDER BY order_id",
List.of("FAILED")));
MutableClock clock = new MutableClock(Instant.parse("2026-07-21T10:00:00Z"));
RunContext context = HarnessCoreFixtures.core(clock).startRun("session", "run-mysql");
MysqlRawResult result = new JdbcMysqlReadOnlyExecutor(Map.of("order_readonly", dataSource), clock)
.execute(plan, context);
assertEquals(List.of("ORDER_ID", "STATUS"), result.columns());
assertEquals(1, result.rows().size());
assertTrue(result.truncated());
}
@Test
void refusesExecutionAfterRunCancellation() throws Exception {
JdbcDataSource dataSource = dataSource();
MysqlDataSourceDefinition definition = definition(MysqlToolLimits.defaults());
MysqlSqlValidator validator = new MysqlSqlValidator(Map.of("order_readonly", definition));
MysqlQueryPlan plan = validator.validate(new MysqlToolRequest(
"order_readonly", "SELECT order_id FROM biz_order", List.of()));
MutableClock clock = new MutableClock(Instant.parse("2026-07-21T10:00:00Z"));
RunContext context = HarnessCoreFixtures.core(clock).startRun("session", "run-cancelled");
context.cancellation().cancel(RunCancellationReason.USER_REQUESTED);
assertThrows(SQLException.class,
() -> new JdbcMysqlReadOnlyExecutor(Map.of("order_readonly", dataSource), clock)
.execute(plan, context));
}
@Test
void appliesReadOnlyTimeoutMaxRowsAndParameterBinding() throws Exception {
javax.sql.DataSource dataSource = mock(javax.sql.DataSource.class);
Connection connection = mock(Connection.class);
PreparedStatement statement = mock(PreparedStatement.class);
ResultSet resultSet = mock(ResultSet.class);
ResultSetMetaData metadata = mock(ResultSetMetaData.class);
when(dataSource.getConnection()).thenReturn(connection);
when(connection.prepareStatement(anyString(), eq(ResultSet.TYPE_FORWARD_ONLY), eq(ResultSet.CONCUR_READ_ONLY)))
.thenReturn(statement);
when(statement.executeQuery()).thenReturn(resultSet);
when(resultSet.getMetaData()).thenReturn(metadata);
when(metadata.getColumnCount()).thenReturn(1);
when(metadata.getColumnLabel(1)).thenReturn("order_id");
when(resultSet.next()).thenReturn(false);
MysqlDataSourceDefinition definition = definition(new MysqlToolLimits(7, 20, 1024, 3));
MysqlQueryPlan plan = new MysqlSqlValidator(Map.of("order_readonly", definition))
.validate(new MysqlToolRequest("order_readonly",
"SELECT order_id FROM biz_order WHERE status = ?", List.of("FAILED")));
MutableClock clock = new MutableClock(Instant.parse("2026-07-21T10:00:00Z"));
RunContext context = HarnessCoreFixtures.core(clock).startRun("session", "run-controls");
new JdbcMysqlReadOnlyExecutor(Map.of("order_readonly", dataSource), clock).execute(plan, context);
verify(connection).setReadOnly(true);
verify(statement).setQueryTimeout(3);
verify(statement).setMaxRows(8);
verify(statement).setObject(1, "FAILED");
}
private JdbcDataSource dataSource() throws Exception {
JdbcDataSource dataSource = new JdbcDataSource();
dataSource.setURL("jdbc:h2:mem:mysqltool;MODE=MySQL;DB_CLOSE_DELAY=-1");
try (var connection = dataSource.getConnection(); var statement = connection.createStatement()) {
statement.execute("DROP TABLE IF EXISTS biz_order");
statement.execute("CREATE TABLE biz_order(order_id VARCHAR(32), status VARCHAR(20))");
statement.execute("INSERT INTO biz_order VALUES ('order-1','FAILED'),('order-2','FAILED'),('order-3','OK')");
}
return dataSource;
}
private MysqlDataSourceDefinition definition(MysqlToolLimits limits) {
return new MysqlDataSourceDefinition("order_readonly", "PUBLIC",
Map.of("PUBLIC", Map.of("biz_order", Set.of("order_id", "status"))), limits);
}
}
@@ -0,0 +1,67 @@
package com.superbiz.agent.harness.tool.mysql;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.contract.EvidenceStatus;
import com.superbiz.agent.harness.tool.boundary.ProjectedToolResult;
import com.superbiz.agent.harness.tool.contract.MysqlToolRequest;
import org.junit.jupiter.api.Test;
import java.util.List;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertTrue;
class MysqlResultProjectorTest {
private final ObjectMapper objectMapper = new ObjectMapper();
@Test
void projectsRowsRedactsSensitiveColumnsAndBoundsValues() throws Exception {
MysqlResultProjector projector = new MysqlResultProjector(objectMapper,
new MysqlToolLimits(1, 5, 1024, 5));
String raw = "{\"columns\":[\"order_id\",\"password\",\"status\"],\"rows\":["
+ "{\"order_id\":\"order-123456\",\"password\":\"secret\",\"status\":\"FAILED\"},"
+ "{\"order_id\":\"order-2\",\"password\":\"secret2\",\"status\":\"OK\"}],\"truncated\":false}";
ProjectedToolResult projected = projector.project(
new MysqlToolRequest("order_readonly", "SELECT order_id FROM biz_order", List.of()),
"framework-mysql-1", raw);
JsonNode json = objectMapper.readTree(projected.agentResult());
assertEquals(EvidenceStatus.EVIDENCE_FOUND, projected.evidenceStatus());
assertEquals(1, json.path("returned_count").asInt());
assertTrue(json.path("truncated").asBoolean());
assertEquals("[REDACTED]", json.path("rows").get(0).path("password").asText());
assertEquals("order", json.path("rows").get(0).path("order_id").asText());
assertFalse(projected.agentResult().contains("secret"));
}
@Test
void returnsNoEvidenceForEmptyRows() throws Exception {
MysqlResultProjector projector = new MysqlResultProjector(objectMapper);
ProjectedToolResult projected = projector.project(
new MysqlToolRequest("order_readonly", "SELECT order_id FROM biz_order", List.of()),
"framework-mysql-2", "{\"columns\":[\"order_id\"],\"rows\":[],\"truncated\":false}");
assertEquals(EvidenceStatus.NO_EVIDENCE, projected.evidenceStatus());
}
@Test
void enforcesTotalUtf8BudgetWithValidJson() throws Exception {
MysqlResultProjector projector = new MysqlResultProjector(objectMapper,
new MysqlToolLimits(10, 100, 230, 5));
String raw = "{\"columns\":[\"order_id\"],\"rows\":["
+ "{\"order_id\":\"" + "a".repeat(80) + "\"},"
+ "{\"order_id\":\"" + "b".repeat(80) + "\"}],\"truncated\":false}";
ProjectedToolResult projected = projector.project(
new MysqlToolRequest("order_readonly", "SELECT order_id FROM biz_order", List.of()),
"framework-mysql-3", raw);
JsonNode json = objectMapper.readTree(projected.agentResult());
assertTrue(json.path("truncated").asBoolean());
assertTrue(projected.agentResult().getBytes(java.nio.charset.StandardCharsets.UTF_8).length <= 230);
}
}
@@ -0,0 +1,68 @@
package com.superbiz.agent.harness.tool.mysql;
import com.superbiz.agent.harness.tool.contract.MysqlToolRequest;
import org.junit.jupiter.api.Test;
import java.util.List;
import java.util.Map;
import java.util.Set;
import static org.junit.jupiter.api.Assertions.assertDoesNotThrow;
import static org.junit.jupiter.api.Assertions.assertThrows;
class MysqlSqlValidatorTest {
private final MysqlSqlValidator validator = new MysqlSqlValidator(Map.of(
"order_readonly", new MysqlDataSourceDefinition(
"order_readonly", "order_db",
Map.of("order_db", Map.of(
"biz_order", Set.of("order_id", "status", "updated_at"),
"payment_record", Set.of("order_id", "status"))),
MysqlToolLimits.defaults())));
@Test
void acceptsExplicitSelectAndCountStar() {
assertDoesNotThrow(() -> validator.validate(new MysqlToolRequest(
"order_readonly",
"SELECT o.order_id, o.status FROM order_db.biz_order o WHERE o.order_id = ?",
List.of("order-1"))));
assertDoesNotThrow(() -> validator.validate(new MysqlToolRequest(
"order_readonly",
"SELECT COUNT(*) AS total FROM order_db.biz_order",
List.of())));
assertDoesNotThrow(() -> validator.validate(new MysqlToolRequest(
"order_readonly",
"SELECT o.order_id, p.status FROM order_db.biz_order o INNER JOIN order_db.payment_record p ON o.order_id = p.order_id WHERE o.status = ?",
List.of("FAILED"))));
}
@Test
void rejectsWritesCompoundQueriesWildcardsAndAllowlistBypass() {
List<String> unsafe = List.of(
"UPDATE order_db.biz_order SET status = ?",
"WITH x AS (SELECT order_id FROM order_db.biz_order) SELECT order_id FROM x",
"SELECT * FROM order_db.biz_order",
"SELECT order_id FROM order_db.biz_order WHERE order_id IN (SELECT order_id FROM order_db.payment_record)",
"SELECT order_id FROM order_db.biz_order UNION SELECT order_id FROM order_db.payment_record",
"SELECT order_id FROM order_db.biz_order WHERE secret = ?",
"SELECT order_id FROM order_db.unknown_table",
"SELECT SLEEP(?) FROM order_db.biz_order",
"SELECT order_id FROM order_db.biz_order WHERE status = 'FAILED'",
"SELECT CASE WHEN status = ? THEN order_id END FROM order_db.biz_order",
"SELECT order_id FROM order_db.biz_order FOR UPDATE",
"SELECT order_id FROM order_db.biz_order; SELECT status FROM order_db.biz_order");
for (String sql : unsafe) {
assertThrows(MysqlSecurityException.class,
() -> validator.validate(new MysqlToolRequest("order_readonly", sql, List.of("x"))), sql);
}
}
@Test
void rejectsPlaceholderMismatchAndAmbiguousColumns() {
assertThrows(MysqlSecurityException.class, () -> validator.validate(new MysqlToolRequest(
"order_readonly", "SELECT order_id FROM order_db.biz_order WHERE order_id = ?", List.of())));
assertThrows(MysqlSecurityException.class, () -> validator.validate(new MysqlToolRequest(
"order_readonly", "SELECT order_id FROM order_db.biz_order o INNER JOIN order_db.payment_record p ON o.order_id = p.order_id", List.of())));
}
}