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: {}
|
||||
|
||||
Reference in New Issue
Block a user