feat(harness): add readonly mysql tool
This commit is contained in:
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
+112
@@ -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())));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user