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,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);
}
}