diff --git a/Dockerfile b/Dockerfile index 2518fb95..5065fe51 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,4 +1,4 @@ -FROM jumpserver/chen-base:20260701_092702 AS stage-build +FROM jumpserver/chen-base:20260811_065334 AS stage-build ENV LANG=en_US.UTF-8 WORKDIR /opt/chen/ diff --git a/Dockerfile-base b/Dockerfile-base index 6b6557db..d2087272 100644 --- a/Dockerfile-base +++ b/Dockerfile-base @@ -24,7 +24,7 @@ RUN set -ex \ && dpkg -i check_linux_${TARGETARCH}.deb \ && rm -f check_linux_${TARGETARCH}.deb -ARG WISP_VERSION=v0.2.13 +ARG WISP_VERSION=v0.2.14 RUN set -ex \ && wget https://github.com/jumpserver/wisp/releases/download/${WISP_VERSION}/wisp_linux_${TARGETARCH}.deb \ && dpkg -i wisp_linux_${TARGETARCH}.deb \ diff --git a/backend/framework/pom.xml b/backend/framework/pom.xml index aaedccc4..16ff34a6 100644 --- a/backend/framework/pom.xml +++ b/backend/framework/pom.xml @@ -19,9 +19,9 @@ - com.alibaba + com.jumpserver druid - 1.2.18 + 1.2.28-jms-chen.1 diff --git a/backend/framework/src/main/java/org/jumpserver/chen/framework/console/QueryConsole.java b/backend/framework/src/main/java/org/jumpserver/chen/framework/console/QueryConsole.java index a3fab598..ab78aa33 100644 --- a/backend/framework/src/main/java/org/jumpserver/chen/framework/console/QueryConsole.java +++ b/backend/framework/src/main/java/org/jumpserver/chen/framework/console/QueryConsole.java @@ -37,9 +37,8 @@ import java.util.LinkedHashMap; import java.util.List; import java.util.Map; -import java.util.concurrent.ConcurrentHashMap; +import java.util.Optional; import java.util.concurrent.CountDownLatch; -import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; @Slf4j @@ -51,6 +50,7 @@ public class QueryConsole extends AbstractConsole { private StateManager stateManager; private final Map dataViews = new HashMap<>(); private volatile Map allowedContexts = Map.of(); + private final SQLChunkTransferManager sqlChunkTransfers = new SQLChunkTransferManager(); private static final Gson GSON = new Gson(); @@ -92,15 +92,13 @@ public void onConnect(Connect connect) { try { var currentContext = this.getSqlActuator().getCurrentSchema(); - if (StringUtils.isEmpty(currentContext) && !StringUtils.isEmpty(context)) { + if (StringUtils.isNotEmpty(context) + && !StringUtils.equals(currentContext, context)) { this.getSqlActuator().changeSchema(context); - this.getState().setCurrentContext(context); - } else { - if (!StringUtils.isEmpty(context) && !currentContext.equals(context)) { - this.getSqlActuator().changeSchema(context); - this.getState().setCurrentContext(context); - } } + this.getState().setCurrentContext( + StringUtils.defaultIfEmpty(context, currentContext) + ); var schemas = this.getSqlActuator().getSchemas(); this.replaceAllowedContexts(schemas); @@ -168,7 +166,7 @@ private void onAction(QueryConsoleAction action) { this.handleSQLChunk(action); } case QueryConsoleAction.ACTION_RUN_SQL_COMPLETE -> { - this.handleSQLComplete(); + this.handleSQLComplete(action); } case QueryConsoleAction.ACTION_RUN_SQL_FILE -> { @@ -184,6 +182,7 @@ private void onAction(QueryConsoleAction action) { case QueryConsoleAction.ACTION_CANCEL -> { + this.sqlChunkTransfers.cancelAll(); this.onCancel(); this.getState().setInQuery(false); this.stateManager.commit(); @@ -195,65 +194,26 @@ private void onAction(QueryConsoleAction action) { } } - private final ConcurrentHashMap sqlChunks = new ConcurrentHashMap<>(); - private CountDownLatch latch; - private int expectedChunks = -1; - private void handleSQLChunk(QueryConsoleAction action) { - var data = (Map) action.getData(); - var chunk = (String) data.get("chunk"); - var index = (Integer) data.get("index"); - var total = (Integer) data.get("total"); - - synchronized (this) { - if (expectedChunks == -1) { - expectedChunks = total; - latch = new CountDownLatch(total); - } - } - - if (sqlChunks.putIfAbsent(index, chunk) == null) { - latch.countDown(); + Optional sql; + try { + sql = this.sqlChunkTransfers.receiveChunk(action.getData()); + } catch (IllegalArgumentException | IllegalStateException e) { + this.getConsoleLogger().error("Invalid SQL chunk transfer: %s", e.getMessage()); + return; } + sql.ifPresent(this::onSQL); } - /** - * 处理分段 SQL 接收完成 - */ - private void handleSQLComplete() { + private void handleSQLComplete(QueryConsoleAction action) { + Optional sql; try { - - // 等待所有分段接收完成 - boolean completed = latch.await(10, TimeUnit.SECONDS); // 超时10秒 - - if (!completed) { - this.getConsoleLogger().error("read sql message timeout!!"); - return; - } - - // 按照索引顺序合并所有分段 - StringBuilder sqlBuilder = new StringBuilder(); - for (int i = 0; i < expectedChunks; i++) { - sqlBuilder.append(sqlChunks.get(i)); - } - - // 合并完成后清理缓存 - var sql = sqlBuilder.toString(); - sqlChunks.clear(); - expectedChunks = -1; - - // 执行完整 SQL - this.getState().setInQuery(true); - this.stateManager.commit(); - - this.onSQL(sql); - - } catch (InterruptedException e) { - Thread.currentThread().interrupt(); - } finally { - this.getState().setInQuery(false); - this.stateManager.commit(); + sql = this.sqlChunkTransfers.receiveComplete(action.getData()); + } catch (IllegalArgumentException | IllegalStateException e) { + this.getConsoleLogger().error("Invalid SQL chunk transfer: %s", e.getMessage()); + return; } + sql.ifPresent(this::onSQL); } private void onDataViewAction(DataViewAction action) { @@ -538,6 +498,7 @@ private void sendDataView(DataView dataView, boolean clearOthers) { @Override public void close() { + this.sqlChunkTransfers.close(); if (this.currentPlan != null) { // flush var session = SessionManager.getCurrentSession(); diff --git a/backend/framework/src/main/java/org/jumpserver/chen/framework/console/SQLChunkTransferManager.java b/backend/framework/src/main/java/org/jumpserver/chen/framework/console/SQLChunkTransferManager.java new file mode 100644 index 00000000..35b4a571 --- /dev/null +++ b/backend/framework/src/main/java/org/jumpserver/chen/framework/console/SQLChunkTransferManager.java @@ -0,0 +1,289 @@ +package org.jumpserver.chen.framework.console; + +import java.math.BigDecimal; +import java.time.Duration; +import java.util.Arrays; +import java.util.HashMap; +import java.util.Map; +import java.util.Objects; +import java.util.Optional; +import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.ScheduledFuture; +import java.util.concurrent.ScheduledThreadPoolExecutor; +import java.util.concurrent.TimeUnit; + +final class SQLChunkTransferManager implements AutoCloseable { + + static final int MAX_CHUNK_SIZE = 4096; + static final int MAX_CHUNKS = 1024; + private static final int MAX_ACTIVE_TRANSFERS = 2; + private static final int MAX_TRACKED_TRANSFERS = 64; + private static final int MAX_REQUEST_ID_LENGTH = 128; + private static final Duration DEFAULT_TIMEOUT = Duration.ofSeconds(30); + private static final ScheduledThreadPoolExecutor TIMEOUT_EXECUTOR = createTimeoutExecutor(); + + private final Object lock = new Object(); + private final Map transfers = new HashMap<>(); + private final ScheduledExecutorService scheduler; + private final Duration timeout; + private boolean closed; + + SQLChunkTransferManager() { + this(TIMEOUT_EXECUTOR, DEFAULT_TIMEOUT); + } + + SQLChunkTransferManager(ScheduledExecutorService scheduler, Duration timeout) { + this.scheduler = Objects.requireNonNull(scheduler, "scheduler"); + this.timeout = Objects.requireNonNull(timeout, "timeout"); + if (timeout.isNegative() || timeout.isZero()) { + throw new IllegalArgumentException("timeout must be positive"); + } + } + + Optional receiveChunk(Object rawData) { + Map data = requireMap(rawData); + String requestId = requireRequestId(data); + int total; + int index; + String chunk; + try { + total = requireTotal(data); + index = requireInt(data, "index"); + if (index < 0 || index >= total) { + throw new IllegalArgumentException("index is outside the transfer range"); + } + chunk = requireChunk(data); + } catch (IllegalArgumentException e) { + reject(requestId); + throw e; + } + + synchronized (lock) { + ensureOpen(); + ChunkTransfer transfer = getOrCreate(requestId, total); + if (transfer.terminal) { + return Optional.empty(); + } + if (transfer.total != total) { + rejectLocked(transfer); + throw new IllegalArgumentException("total does not match the existing transfer"); + } + + String existing = transfer.chunks[index]; + if (existing != null) { + if (!existing.equals(chunk)) { + rejectLocked(transfer); + throw new IllegalArgumentException("duplicate chunk has different content"); + } + return Optional.empty(); + } + + transfer.chunks[index] = chunk; + transfer.receivedChunks += 1; + return assembleIfReady(transfer); + } + } + + Optional receiveComplete(Object rawData) { + Map data = requireMap(rawData); + String requestId = requireRequestId(data); + int total; + try { + total = requireTotal(data); + } catch (IllegalArgumentException e) { + reject(requestId); + throw e; + } + + synchronized (lock) { + ensureOpen(); + ChunkTransfer transfer = getOrCreate(requestId, total); + if (transfer.terminal) { + return Optional.empty(); + } + if (transfer.total != total) { + rejectLocked(transfer); + throw new IllegalArgumentException("total does not match the existing transfer"); + } + + transfer.completeReceived = true; + return assembleIfReady(transfer); + } + } + + void cancelAll() { + synchronized (lock) { + clearTransfersLocked(); + } + } + + @Override + public void close() { + synchronized (lock) { + if (closed) { + return; + } + closed = true; + clearTransfersLocked(); + } + } + + private ChunkTransfer getOrCreate(String requestId, int total) { + ChunkTransfer existing = transfers.get(requestId); + if (existing != null) { + return existing; + } + if (transfers.size() >= MAX_TRACKED_TRANSFERS) { + throw new IllegalArgumentException("too many tracked chunk transfers"); + } + long activeTransfers = transfers.values().stream() + .filter(transfer -> !transfer.terminal) + .count(); + if (activeTransfers >= MAX_ACTIVE_TRANSFERS) { + throw new IllegalArgumentException("too many active chunk transfers"); + } + + ChunkTransfer transfer = new ChunkTransfer(total); + transfers.put(requestId, transfer); + try { + transfer.timeoutFuture = scheduler.schedule( + () -> expire(requestId, transfer), + timeout.toMillis(), + TimeUnit.MILLISECONDS + ); + } catch (RuntimeException e) { + transfers.remove(requestId, transfer); + throw e; + } + return transfer; + } + + private Optional assembleIfReady(ChunkTransfer transfer) { + if (!transfer.completeReceived || transfer.receivedChunks != transfer.total) { + return Optional.empty(); + } + + StringBuilder sql = new StringBuilder(); + for (String chunk : transfer.chunks) { + sql.append(chunk); + } + String assembled = sql.toString(); + rejectLocked(transfer); + return Optional.of(assembled); + } + + private void reject(String requestId) { + synchronized (lock) { + ChunkTransfer transfer = transfers.get(requestId); + if (transfer != null) { + rejectLocked(transfer); + } + } + } + + private void rejectLocked(ChunkTransfer transfer) { + if (transfer.terminal) { + return; + } + transfer.terminal = true; + transfer.completeReceived = false; + transfer.receivedChunks = 0; + Arrays.fill(transfer.chunks, null); + } + + private void expire(String requestId, ChunkTransfer expectedTransfer) { + synchronized (lock) { + if (transfers.remove(requestId, expectedTransfer)) { + rejectLocked(expectedTransfer); + } + } + } + + private void clearTransfersLocked() { + for (ChunkTransfer transfer : transfers.values()) { + rejectLocked(transfer); + if (transfer.timeoutFuture != null) { + transfer.timeoutFuture.cancel(false); + } + } + transfers.clear(); + } + + private void ensureOpen() { + if (closed) { + throw new IllegalStateException("chunk transfer manager is closed"); + } + } + + private static Map requireMap(Object data) { + if (!(data instanceof Map map)) { + throw new IllegalArgumentException("chunk data must be an object"); + } + return map; + } + + private static String requireRequestId(Map data) { + Object value = data.get("requestId"); + if (!(value instanceof String requestId) + || requestId.isBlank() + || requestId.length() > MAX_REQUEST_ID_LENGTH) { + throw new IllegalArgumentException("requestId is invalid"); + } + return requestId; + } + + private static int requireTotal(Map data) { + int total = requireInt(data, "total"); + if (total <= 0 || total > MAX_CHUNKS) { + throw new IllegalArgumentException("total is outside the allowed range"); + } + return total; + } + + private static int requireInt(Map data, String key) { + Object value = data.get(key); + if (!(value instanceof Number number)) { + throw new IllegalArgumentException(key + " must be a number"); + } + try { + return new BigDecimal(number.toString()).intValueExact(); + } catch (ArithmeticException | NumberFormatException e) { + throw new IllegalArgumentException(key + " must be a finite int", e); + } + } + + private static String requireChunk(Map data) { + Object value = data.get("chunk"); + if (!(value instanceof String chunk)) { + throw new IllegalArgumentException("chunk must be a string"); + } + if (chunk.length() > MAX_CHUNK_SIZE) { + throw new IllegalArgumentException("chunk exceeds the allowed size"); + } + return chunk; + } + + private static ScheduledThreadPoolExecutor createTimeoutExecutor() { + ScheduledThreadPoolExecutor executor = new ScheduledThreadPoolExecutor(1, runnable -> { + Thread thread = new Thread(runnable, "query-console-chunk-timeout"); + thread.setDaemon(true); + return thread; + }); + executor.setRemoveOnCancelPolicy(true); + return executor; + } + + private static final class ChunkTransfer { + private final int total; + private final String[] chunks; + private int receivedChunks; + private boolean completeReceived; + private boolean terminal; + private ScheduledFuture timeoutFuture; + + private ChunkTransfer(int total) { + this.total = total; + this.chunks = new String[total]; + } + } +} diff --git a/backend/framework/src/main/java/org/jumpserver/chen/framework/console/dataview/export/DataExport.java b/backend/framework/src/main/java/org/jumpserver/chen/framework/console/dataview/export/DataExport.java index 8b3a10b0..131757d6 100644 --- a/backend/framework/src/main/java/org/jumpserver/chen/framework/console/dataview/export/DataExport.java +++ b/backend/framework/src/main/java/org/jumpserver/chen/framework/console/dataview/export/DataExport.java @@ -141,6 +141,7 @@ public void exportData(String path, DataViewData data) throws Exception { } else if (obj instanceof Date) { SimpleDateFormat fmt = new SimpleDateFormat("yyyy-MM-dd HH:mm:ss"); writeString(writer, fmt.format(obj)); + writer.write(","); } else { writeString(writer, row.get(field.getName())); writer.write(","); @@ -157,9 +158,9 @@ public void exportData(String path, DataViewData data) throws Exception { private static void writeString(BufferedWriter writer, Object object) throws IOException { var str = object.toString(); - if (str.contains(",")) { - str = "\"" + str + "\""; + if (str.contains("\"") || str.contains(",")) { + str = "\"" + str.replace("\"", "\"\"") + "\""; } writer.write(str); } -} \ No newline at end of file +} diff --git a/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseActionHandler.java b/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseActionHandler.java index a6136410..f2fa7bd7 100644 --- a/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseActionHandler.java +++ b/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseActionHandler.java @@ -162,7 +162,9 @@ public EventEmitter onShowProperties(TreeNode node) { public EventEmitter onShowObjectProperties(String type, String sql, TreeNode node) throws SQLException { var sqlActuator = this.getDatasource().getConnectionManager().getSqlActuator(); var objName = TreeUtils.getValue(node.getKey(), type); - var result = sqlActuator.execute(SQL.of(sql, objName)); + var command = SQL.of(sql, objName.replace("'", "''")); + log.info("resource action show_properties: type={}, node={}", type, node.getKey()); + var result = sqlActuator.execute(command); var detailDialog = new DetailDialog(node.getKey(), type + MessageUtils.get("Properties")); detailDialog.setWidth("50%"); diff --git a/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseConnectionManager.java b/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseConnectionManager.java index 577a3699..b3d1dfce 100644 --- a/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseConnectionManager.java +++ b/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseConnectionManager.java @@ -10,6 +10,7 @@ import org.jumpserver.chen.framework.driver.DriverClassLoader; import org.jumpserver.chen.framework.driver.DriverManager; import org.jumpserver.chen.framework.i18n.MessageUtils; +import org.jumpserver.chen.framework.utils.SqlIdentifierUtils; import java.lang.reflect.InvocationTargetException; import java.sql.Connection; @@ -120,6 +121,9 @@ public DruidDataSource getOrInitDataSource(String database) throws SQLException if (StringUtils.isEmpty(database)) { database = this.connectInfo.getDb(); } + // Reject URL metacharacters before interpolation into the JDBC URL, + // preventing injection of driver connection properties. + SqlIdentifierUtils.validateDatabaseName(database); if (this.dataSourceMap.containsKey(database)) { return this.dataSourceMap.get(database); } diff --git a/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseSQLActuator.java b/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseSQLActuator.java index 757aa54f..99b0c187 100644 --- a/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseSQLActuator.java +++ b/backend/framework/src/main/java/org/jumpserver/chen/framework/datasource/base/BaseSQLActuator.java @@ -25,6 +25,17 @@ import java.math.BigDecimal; import java.math.BigInteger; import java.sql.*; +import java.time.Instant; +import java.time.LocalDate; +import java.time.LocalDateTime; +import java.time.LocalTime; +import java.time.OffsetDateTime; +import java.time.OffsetTime; +import java.time.ZoneOffset; +import java.time.ZonedDateTime; +import java.time.format.DateTimeFormatter; +import java.time.format.DateTimeFormatterBuilder; +import java.time.temporal.ChronoField; import java.util.ArrayList; import java.util.HashMap; import java.util.List; @@ -41,6 +52,22 @@ public abstract class BaseSQLActuator implements SQLActuator { // Keep large JDBC text values bounded so one cell cannot fail or stall the whole result view. private static final int MAX_TEXT_DISPLAY_LENGTH = 1024 * 1024; private static final String TRUNCATED_SUFFIX = "...[truncated]"; + private static final DateTimeFormatter LOCAL_DATE_TIME_DISPLAY_FORMATTER = new DateTimeFormatterBuilder() + .appendPattern("uuuu-MM-dd HH:mm:ss") + .appendFraction(ChronoField.NANO_OF_SECOND, 0, 9, true) + .toFormatter(); + private static final DateTimeFormatter LOCAL_TIME_DISPLAY_FORMATTER = new DateTimeFormatterBuilder() + .appendPattern("HH:mm:ss") + .appendFraction(ChronoField.NANO_OF_SECOND, 0, 9, true) + .toFormatter(); + private static final DateTimeFormatter OFFSET_DATE_TIME_DISPLAY_FORMATTER = new DateTimeFormatterBuilder() + .append(LOCAL_DATE_TIME_DISPLAY_FORMATTER) + .appendOffsetId() + .toFormatter(); + private static final DateTimeFormatter OFFSET_TIME_DISPLAY_FORMATTER = new DateTimeFormatterBuilder() + .append(LOCAL_TIME_DISPLAY_FORMATTER) + .appendOffsetId() + .toFormatter(); private ConnectionManager connectionManager; private Connection connection; @@ -172,7 +199,7 @@ private void executeStatement(SQLExecutePlan plan, Statement statement, SQLQuery List fs = new ArrayList<>(); for (int i = 1; i <= columnCount; i++) { try { - fs.add(this.normalizeJdbcValue(resultSet.getObject(i))); + fs.add(this.normalizeJdbcValue(resultSet, i)); } catch (NoClassDefFoundError e) { log.error(e.getMessage()); } @@ -202,17 +229,30 @@ private void executeStatement(SQLExecutePlan plan, Statement statement, SQLQuery } } - // Normalize JDBC driver objects before FastJSON sees them in update_data_view packets. + // Normalize JDBC driver objects before Gson sees them in update_data_view packets. + protected Object normalizeJdbcValue(ResultSet resultSet, int columnIndex) throws SQLException { + return this.normalizeJdbcValue(resultSet.getObject(columnIndex)); + } + protected Object normalizeJdbcValue(Object value) throws SQLException { if (value == null) { return null; } - if (value instanceof Timestamp timestamp) { - return new Date(timestamp.getTime()); + if (value.getClass().getName().equals("oracle.sql.TIMESTAMPTZ")) { + return this.normalizeOracleTimestampWithTimeZone(value); + } + + var temporalValue = this.formatTemporalValue(value); + if (temporalValue != null) { + return temporalValue; + } + + if (value instanceof BigDecimal decimal) { + return decimal.toPlainString(); } - if (value instanceof Long || value instanceof BigDecimal || value instanceof BigInteger) { + if (value instanceof Long || value instanceof BigInteger) { return value.toString(); } @@ -240,9 +280,64 @@ protected Object normalizeJdbcValue(Object value) throws SQLException { return value.toString(); } + if (value.getClass().getName().equals("microsoft.sql.DateTimeOffset")) { + return value.toString(); + } + return value; } + private String normalizeOracleTimestampWithTimeZone(Object value) throws SQLException { + try { + var offsetDateTime = value.getClass().getMethod("toOffsetDateTime").invoke(value); + return OFFSET_DATE_TIME_DISPLAY_FORMATTER.format((OffsetDateTime) offsetDateTime); + } catch (ReflectiveOperationException | ClassCastException e) { + var cause = e.getCause() == null ? e : e.getCause(); + if (cause instanceof SQLException sqlException) { + throw sqlException; + } + throw new SQLException("normalize Oracle TIMESTAMP WITH TIME ZONE failed", cause); + } + } + + private String formatTemporalValue(Object value) { + if (value instanceof Timestamp timestamp) { + return LOCAL_DATE_TIME_DISPLAY_FORMATTER.format(timestamp.toLocalDateTime()); + } + if (value instanceof Date date) { + return date.toLocalDate().toString(); + } + if (value instanceof Time time) { + return LOCAL_TIME_DISPLAY_FORMATTER.format(time.toLocalTime()); + } + if (value instanceof LocalDateTime localDateTime) { + return LOCAL_DATE_TIME_DISPLAY_FORMATTER.format(localDateTime); + } + if (value instanceof LocalDate localDate) { + return localDate.toString(); + } + if (value instanceof LocalTime localTime) { + return LOCAL_TIME_DISPLAY_FORMATTER.format(localTime); + } + if (value instanceof OffsetDateTime offsetDateTime) { + return OFFSET_DATE_TIME_DISPLAY_FORMATTER.format(offsetDateTime); + } + if (value instanceof OffsetTime offsetTime) { + return OFFSET_TIME_DISPLAY_FORMATTER.format(offsetTime); + } + if (value instanceof ZonedDateTime zonedDateTime) { + var formatted = OFFSET_DATE_TIME_DISPLAY_FORMATTER.format(zonedDateTime); + if (!(zonedDateTime.getZone() instanceof ZoneOffset)) { + formatted += "[" + zonedDateTime.getZone().getId() + "]"; + } + return formatted; + } + if (value instanceof Instant instant) { + return instant.toString(); + } + return null; + } + private String toDisplayArray(java.sql.Array jdbcArray) throws SQLException { try { var text = jdbcArray.toString(); diff --git a/backend/framework/src/main/java/org/jumpserver/chen/framework/session/SessionManager.java b/backend/framework/src/main/java/org/jumpserver/chen/framework/session/SessionManager.java index 6778a070..659ad02b 100644 --- a/backend/framework/src/main/java/org/jumpserver/chen/framework/session/SessionManager.java +++ b/backend/framework/src/main/java/org/jumpserver/chen/framework/session/SessionManager.java @@ -9,9 +9,13 @@ @Slf4j public class SessionManager { + // 绑定创建 Chen 会话的 Servlet HTTP session,WS 握手时用它阻止 token 被跨浏览器重放。 + public static final String WEB_SESSION_ID_ATTRIBUTE = "webSessionId"; private final static SessionManager instance = new SessionManager(); private final static ThreadLocal token = new ThreadLocal<>(); private final Map store = new ConcurrentHashMap<>(); + // 每个 Chen 会话只能有一个主 /ws/session 连接,值为 WebSocket session id。 + private final Map primaryWebSockets = new ConcurrentHashMap<>(); public static String registerSession(Session session) { String token = createToken(); @@ -23,6 +27,7 @@ public static String registerSession(Session session) { public static void unregisterSession(String token) { instance.store.remove(token); + instance.primaryWebSockets.remove(token); log.info("session {} unregistered, current session count {}", token, instance.getCurrentSessionCount()); } @@ -54,6 +59,17 @@ public static Session getSession(String token) { return instance.store.get(token); } + public static boolean claimPrimaryWebSocket(String token, String webSocketId) { + // 原子占用,避免并发握手同时替换当前主连接。 + String existing = instance.primaryWebSockets.putIfAbsent(token, webSocketId); + return existing == null || existing.equals(webSocketId); + } + + public static boolean releasePrimaryWebSocket(String token, String webSocketId) { + // 只允许占用者释放,防止被拒绝的重放连接关闭正常会话。 + return instance.primaryWebSockets.remove(token, webSocketId); + } + private static String createToken() { return UUID.randomUUID().toString().replace("-", ""); diff --git a/backend/framework/src/main/java/org/jumpserver/chen/framework/utils/PageUtils.java b/backend/framework/src/main/java/org/jumpserver/chen/framework/utils/PageUtils.java index ad24e8b4..b6cb9631 100644 --- a/backend/framework/src/main/java/org/jumpserver/chen/framework/utils/PageUtils.java +++ b/backend/framework/src/main/java/org/jumpserver/chen/framework/utils/PageUtils.java @@ -15,7 +15,6 @@ import com.alibaba.druid.sql.dialect.oracle.visitor.OracleASTVisitorAdapter; import com.alibaba.druid.sql.dialect.postgresql.ast.stmt.PGSelectQueryBlock; import com.alibaba.druid.sql.dialect.sqlserver.ast.SQLServerSelectQueryBlock; -import com.alibaba.druid.sql.dialect.sqlserver.ast.SQLServerTop; import com.alibaba.druid.util.JdbcUtils; import java.util.Iterator; @@ -264,14 +263,14 @@ private static boolean limitSQLServer(SQLSelect select, DbType dbType, int offse if (query instanceof SQLSelectQueryBlock) { queryBlock = (SQLServerSelectQueryBlock) query; if (offset <= 0) { - SQLServerTop top = queryBlock.getTop(); + SQLTop top = queryBlock.getTop(); if (check && top != null && !top.isPercent() && top.getExpr() instanceof SQLNumericLiteralExpr) { int rowCount = ((SQLNumericLiteralExpr) top.getExpr()).getNumber().intValue(); if (rowCount <= count) { return false; } } - queryBlock.setTop(new SQLServerTop(new SQLNumberExpr(count))); + queryBlock.setTop(new SQLTop(new SQLNumberExpr(count))); return true; } else { // 创建 SELECT NULL 的子查询 @@ -303,7 +302,7 @@ private static boolean limitSQLServer(SQLSelect select, DbType dbType, int offse } else { queryBlock = new SQLServerSelectQueryBlock(); if (offset <= 0) { - queryBlock.setTop(new SQLServerTop(new SQLNumberExpr(count))); + queryBlock.setTop(new SQLTop(new SQLNumberExpr(count))); select.setQuery(queryBlock); return true; } else { diff --git a/backend/framework/src/main/java/org/jumpserver/chen/framework/utils/SqlIdentifierUtils.java b/backend/framework/src/main/java/org/jumpserver/chen/framework/utils/SqlIdentifierUtils.java new file mode 100644 index 00000000..a68270eb --- /dev/null +++ b/backend/framework/src/main/java/org/jumpserver/chen/framework/utils/SqlIdentifierUtils.java @@ -0,0 +1,40 @@ +package org.jumpserver.chen.framework.utils; + +import lombok.extern.slf4j.Slf4j; + +import java.sql.SQLException; + +/** + * Validates runtime database identifiers before they are interpolated into a + * JDBC URL, so that URL/property metacharacters in a database name cannot + * inject driver connection properties (e.g. PostgreSQL socketFactory/loggerFile, + * SQLServer ";prop=val"). + */ +@Slf4j +public final class SqlIdentifierUtils { + + // Characters that act as delimiters in at least one supported JDBC URL form. + // Rejecting them keeps a database name from escaping its ${db} placeholder. + private static final String FORBIDDEN_CHARS = "?&/:@#\\;= \t\r\n"; + + private SqlIdentifierUtils() { + } + + /** + * Reject database names containing URL/property metacharacters or control + * characters. Blank values are allowed; callers fall back to the configured db. + */ + public static void validateDatabaseName(String name) throws SQLException { + if (name == null || name.isEmpty()) { + return; + } + for (int i = 0; i < name.length(); i++) { + char c = name.charAt(i); + if (c < 0x20 || FORBIDDEN_CHARS.indexOf(c) >= 0) { + // Do not echo the value: it may contain log-forging characters. + log.warn("Rejected database name containing URL metacharacter"); + throw new SQLException("Invalid database identifier"); + } + } + } +} diff --git a/backend/framework/src/main/java/org/jumpserver/chen/framework/ws/SessionWebSocketHandler.java b/backend/framework/src/main/java/org/jumpserver/chen/framework/ws/SessionWebSocketHandler.java index 1c0b2a0f..3b56263b 100644 --- a/backend/framework/src/main/java/org/jumpserver/chen/framework/ws/SessionWebSocketHandler.java +++ b/backend/framework/src/main/java/org/jumpserver/chen/framework/ws/SessionWebSocketHandler.java @@ -30,6 +30,13 @@ public void afterConnectionEstablished(WebSocketSession session) throws Exceptio } var token = (String) session.getAttributes().get("token"); + // 拒绝第二个主会话,避免重放连接替换 PacketIO 并在关闭时终止原会话。 + if (!SessionManager.claimPrimaryWebSocket(token, session.getId())) { + log.warn("Reject duplicate primary WebSocket connection"); + session.close(CloseStatus.POLICY_VIOLATION); + return; + } + log.info("Primary WebSocket connection established"); SessionManager.setContext(token); Session sess = SessionManager.getCurrentSession(); @@ -96,10 +103,15 @@ public void handleMessage(WebSocketSession session, WebSocketMessage message) @Override public void afterConnectionClosed(WebSocketSession session, CloseStatus closeStatus) throws Exception { var token = (String) session.getAttributes().get("token"); + // 被拒绝的连接没有占用主会话,不得影响当前仍在线的用户。 + if (!SessionManager.releasePrimaryWebSocket(token, session.getId())) { + return; + } + log.info("Primary WebSocket connection closed: code={}", closeStatus.getCode()); SessionManager.setContext(token); var sess = SessionManager.getCurrentSession(); if (sess != null) { sess.close(); } } -} \ No newline at end of file +} diff --git a/backend/modules/src/main/java/org.jumpserver.chen.modules/mariadb/MariaDBActuator.java b/backend/modules/src/main/java/org.jumpserver.chen.modules/mariadb/MariaDBActuator.java index 3ecb8b88..01d6871e 100644 --- a/backend/modules/src/main/java/org.jumpserver.chen.modules/mariadb/MariaDBActuator.java +++ b/backend/modules/src/main/java/org.jumpserver.chen.modules/mariadb/MariaDBActuator.java @@ -4,6 +4,8 @@ import org.jumpserver.chen.modules.mysql.MysqlActuator; import java.sql.Connection; +import java.sql.ResultSet; +import java.sql.SQLException; public class MariaDBActuator extends MysqlActuator { public MariaDBActuator(ConnectionManager connectionManager) { @@ -12,4 +14,13 @@ public MariaDBActuator(ConnectionManager connectionManager) { public MariaDBActuator(MariaDBActuator sqlActuator, Connection connection) { super(sqlActuator, connection); } + + @Override + protected Object normalizeJdbcValue(ResultSet resultSet, int columnIndex) throws SQLException { + var columnTypeName = resultSet.getMetaData().getColumnTypeName(columnIndex); + if ("JSON".equalsIgnoreCase(columnTypeName) || "MYSQL_JSON".equalsIgnoreCase(columnTypeName)) { + return resultSet.getString(columnIndex); + } + return super.normalizeJdbcValue(resultSet, columnIndex); + } } diff --git a/backend/modules/src/main/java/org.jumpserver.chen.modules/mysql/MysqlConnectionManager.java b/backend/modules/src/main/java/org.jumpserver.chen.modules/mysql/MysqlConnectionManager.java index 8d08025d..cc83f291 100644 --- a/backend/modules/src/main/java/org.jumpserver.chen.modules/mysql/MysqlConnectionManager.java +++ b/backend/modules/src/main/java/org.jumpserver.chen.modules/mysql/MysqlConnectionManager.java @@ -17,7 +17,7 @@ public class MysqlConnectionManager extends BaseConnectionManager { - private static final String jdbcUrlTemplate = "jdbc:mysql://${host}:${port}/${db}?useSSL=false&useUnicode=true&characterEncoding=UTF-8&zeroDateTimeBehavior=CONVERT_TO_NULL&tinyInt1isBit=false&jdbcCompliantTruncation=false"; + private static final String jdbcUrlTemplate = "jdbc:mysql://${host}:${port}/${db}?useSSL=false&useUnicode=true&characterEncoding=UTF-8&zeroDateTimeBehavior=CONVERT_TO_NULL&tinyInt1isBit=false&jdbcCompliantTruncation=false&allowPublicKeyRetrieval=true"; private String jdbcUrl; public MysqlConnectionManager(DBConnectInfo connectInfo, Datasource datasource) { diff --git a/backend/modules/src/main/java/org.jumpserver.chen.modules/oracle/OracleActuator.java b/backend/modules/src/main/java/org.jumpserver.chen.modules/oracle/OracleActuator.java index 7912445a..8f18190c 100644 --- a/backend/modules/src/main/java/org.jumpserver.chen.modules/oracle/OracleActuator.java +++ b/backend/modules/src/main/java/org.jumpserver.chen.modules/oracle/OracleActuator.java @@ -1,7 +1,7 @@ package org.jumpserver.chen.modules.oracle; -import com.alibaba.druid.DbType; import com.alibaba.druid.sql.SQLUtils; +import com.alibaba.druid.sql.parser.ParserException; import org.jumpserver.chen.framework.datasource.ConnectionManager; import org.jumpserver.chen.framework.datasource.base.BaseSQLActuator; import org.jumpserver.chen.framework.datasource.sql.SQL; @@ -66,8 +66,15 @@ public SQLExecutePlan createPlan(SQL sql) throws SQLException { @Override public List parseSQL(SQL sql) { - return SQLUtils.parseStatements(sql.getSql(), DbType.ali_oracle).stream() - .map(stmt -> SQLUtils.toSQLString(stmt, DbType.ali_oracle)) + var dbType = this.getDbType(); + var statements = SQLUtils.parseStatements(sql.getSql(), dbType); + for (var i = 0; i < statements.size() - 1; i++) { + if (!statements.get(i).isAfterSemi()) { + throw new ParserException("Multiple SQL statements must be separated by semicolons"); + } + } + return statements.stream() + .map(stmt -> SQLUtils.toSQLString(stmt, dbType)) .toList(); } } diff --git a/backend/modules/src/main/java/org.jumpserver.chen.modules/oracle/OracleConnectionManager.java b/backend/modules/src/main/java/org.jumpserver.chen.modules/oracle/OracleConnectionManager.java index 98dde0e2..26c39ac9 100644 --- a/backend/modules/src/main/java/org.jumpserver.chen.modules/oracle/OracleConnectionManager.java +++ b/backend/modules/src/main/java/org.jumpserver.chen.modules/oracle/OracleConnectionManager.java @@ -6,6 +6,7 @@ import org.jumpserver.chen.framework.datasource.sql.SQL; import java.sql.SQLException; +import java.util.Objects; import java.util.Properties; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ExecutionException; @@ -40,31 +41,23 @@ public void ping() { var pool = Executors.newFixedThreadPool(2); - CompletableFuture f1 = CompletableFuture.supplyAsync(() -> { - try { - this.ping(sidUrl, props); - return sidUrl; - } catch (SQLException e) { - return null; - } - }, pool); - - CompletableFuture f2 = CompletableFuture.supplyAsync(() -> { - try { - this.ping(serviceUrl, props); - return serviceUrl; - } catch (SQLException e) { - return null; - } - }, pool); + CompletableFuture f1 = CompletableFuture.supplyAsync( + () -> this.attemptConnection(sidUrl, props), pool + ); + + CompletableFuture f2 = CompletableFuture.supplyAsync( + () -> this.attemptConnection(serviceUrl, props), pool + ); CompletableFuture combinedFuture = CompletableFuture.allOf(f1, f2); try { combinedFuture.get(); // 等待所有Future完成 // 判断哪个Future成功完成,并设置jdbcUrl - this.jdbcUrl = f1.get() != null ? f1.get() : f2.get(); + var sidAttempt = f1.get(); + var serviceNameAttempt = f2.get(); + this.jdbcUrl = sidAttempt.jdbcUrl() != null ? sidAttempt.jdbcUrl() : serviceNameAttempt.jdbcUrl(); if (this.jdbcUrl == null) { - throw new RuntimeException("Both SID and ServiceName connections failed."); + throw connectionFailed(sidAttempt.error(), serviceNameAttempt.error()); } } catch (InterruptedException | ExecutionException e) { throw new RuntimeException("Error occurred while pinging database", e); @@ -73,6 +66,26 @@ public void ping() { } } + private ConnectionAttempt attemptConnection(String jdbcUrl, Properties props) { + try { + this.ping(jdbcUrl, props); + return new ConnectionAttempt(jdbcUrl, null); + } catch (SQLException e) { + return new ConnectionAttempt(null, e); + } + } + + private RuntimeException connectionFailed(SQLException sidError, SQLException serviceNameError) { + var message = Objects.equals(sidError.getMessage(), serviceNameError.getMessage()) + ? sidError.getMessage() + : "SID: %s; ServiceName: %s".formatted(sidError.getMessage(), serviceNameError.getMessage()); + var error = new RuntimeException(message, sidError); + error.addSuppressed(serviceNameError); + return error; + } + + private record ConnectionAttempt(String jdbcUrl, SQLException error) {} + private static final String SQL_GET_VERSION = "select concat(product,concat(version,status)) as version from product_component_version where product like 'Oracle%'"; @Override diff --git a/backend/modules/src/main/java/org.jumpserver.chen.modules/postgresql/PostgresqlSQLHintsHandler.java b/backend/modules/src/main/java/org.jumpserver.chen.modules/postgresql/PostgresqlSQLHintsHandler.java index 38be4baa..c0b970dc 100644 --- a/backend/modules/src/main/java/org.jumpserver.chen.modules/postgresql/PostgresqlSQLHintsHandler.java +++ b/backend/modules/src/main/java/org.jumpserver.chen.modules/postgresql/PostgresqlSQLHintsHandler.java @@ -6,6 +6,7 @@ import org.jumpserver.chen.framework.datasource.entity.resource.Field; import org.jumpserver.chen.framework.datasource.entity.resource.Table; import org.jumpserver.chen.framework.datasource.sql.SQL; +import org.jumpserver.chen.framework.utils.SqlIdentifierUtils; import org.jumpserver.chen.framework.utils.TreeUtils; import java.sql.SQLException; @@ -46,6 +47,9 @@ public Map> getHints(String nodeKey, String context) throws var db = TreeUtils.getValue(nodeKey, "database"); if (StringUtils.isNotEmpty(db)) { + // nodeKey is client-controlled; reject URL metacharacters before it + // reaches the JDBC URL via setDatabaseContext. + SqlIdentifierUtils.validateDatabaseName(db); this.connectionManager.setDatabaseContext(db); } diff --git a/backend/modules/src/main/java/org.jumpserver.chen.modules/sqlserver/SQLServerConnectionManager.java b/backend/modules/src/main/java/org.jumpserver.chen.modules/sqlserver/SQLServerConnectionManager.java index ab52146e..d5442437 100644 --- a/backend/modules/src/main/java/org.jumpserver.chen.modules/sqlserver/SQLServerConnectionManager.java +++ b/backend/modules/src/main/java/org.jumpserver.chen.modules/sqlserver/SQLServerConnectionManager.java @@ -15,7 +15,7 @@ @Slf4j public class SQLServerConnectionManager extends BaseConnectionManager { - private static final String jdbcUrlTemplate = "jdbc:sqlserver://${host}:${port};DatabaseName=${db};trustServerCertificate=true;"; + private static final String jdbcUrlTemplate = "jdbc:sqlserver://${host}:${port};DatabaseName=${db};encrypt=${encrypt};trustServerCertificate=true;"; private String jdbcUrl; private String driverClassloaderName = "mssql-jdbc-12.10.2.jre11.jar"; @@ -67,7 +67,7 @@ public Driver getDriver() { @Override public void ping() throws SQLException { - var url = this.getConnectInfo().toJDBCUrl(jdbcUrlTemplate); + var url = this.toJDBCUrl(this.getConnectInfo().getDb()); this.ping(url); this.jdbcUrl = url; } @@ -89,12 +89,22 @@ public String getJDBCUrl() { @Override public String getDisplayJDBCUrl() { - return this.getConnectInfo().toDisplayJDBCUrl(jdbcUrlTemplate); + return this.getConnectInfo().toDisplayJDBCUrl(this.getJDBCUrlTemplate()); } @Override public String getJDBCUrl(String database) { - return this.getConnectInfo().toJDBCUrl(jdbcUrlTemplate, database); + return this.toJDBCUrl(database); + } + + private String toJDBCUrl(String database) { + return this.getConnectInfo().toJDBCUrl(this.getJDBCUrlTemplate(), database); + } + + private String getJDBCUrlTemplate() { + var encrypt = this.getConnectInfo().getOptions().get("encrypt"); + var encryptEnabled = encrypt == null || !"false".equalsIgnoreCase(encrypt.toString()); + return jdbcUrlTemplate.replace("${encrypt}", Boolean.toString(encryptEnabled)); } } diff --git a/backend/web/src/main/java/org/jumpserver/chen/web/config/WebSocketConfig.java b/backend/web/src/main/java/org/jumpserver/chen/web/config/WebSocketConfig.java index 690145c2..d2bdc441 100644 --- a/backend/web/src/main/java/org/jumpserver/chen/web/config/WebSocketConfig.java +++ b/backend/web/src/main/java/org/jumpserver/chen/web/config/WebSocketConfig.java @@ -1,6 +1,7 @@ package org.jumpserver.chen.web.config; import lombok.extern.slf4j.Slf4j; +import org.jumpserver.chen.framework.session.SessionManager; import org.jumpserver.chen.framework.ws.ConsoleWebSocketHandler; import org.jumpserver.chen.framework.ws.DBConsoleWebsocketHandler; import org.jumpserver.chen.framework.ws.SessionWebSocketHandler; @@ -9,12 +10,12 @@ import org.springframework.http.HttpStatus; import org.springframework.http.server.ServerHttpRequest; import org.springframework.http.server.ServerHttpResponse; +import org.springframework.http.server.ServletServerHttpRequest; import org.springframework.web.socket.WebSocketHandler; -import org.springframework.web.socket.config.annotation.EnableWebSocket; -import org.springframework.web.socket.config.annotation.WebSocketConfigurer; -import org.springframework.web.socket.config.annotation.WebSocketHandlerRegistry; import org.springframework.web.socket.server.HandshakeInterceptor; import org.springframework.web.socket.server.standard.ServletServerContainerFactoryBean; +import org.springframework.web.socket.server.support.WebSocketHandlerMapping; +import org.springframework.web.socket.server.support.WebSocketHttpRequestHandler; import java.net.InetSocketAddress; import java.net.URI; @@ -22,30 +23,39 @@ import java.util.stream.Collectors; @Configuration -@EnableWebSocket @Slf4j -public class WebSocketConfig implements WebSocketConfigurer { +public class WebSocketConfig { @Bean public ServletServerContainerFactoryBean createWebSocketContainer() { ServletServerContainerFactoryBean container = new ServletServerContainerFactoryBean(); - container.setMaxTextMessageBufferSize(20 * 1024 * 1024); - container.setMaxBinaryMessageBufferSize(20 * 1024 * 1024); + container.setMaxTextMessageBufferSize(1024 * 1024); // 可选:异步发送超时 // container.setAsyncSendTimeout(20_000L); return container; } - @Override - public void registerWebSocketHandlers(WebSocketHandlerRegistry registry) { - registry - .addHandler(new ConsoleWebSocketHandler(), "/ws/console") - .addHandler(new SessionWebSocketHandler(), "/ws/session") - .addHandler(new DBConsoleWebsocketHandler(), "/ws/db-console") - .addInterceptors(new ServletWebSocketHandshakeInterceptor()) - .setAllowedOrigins("*"); + @Bean + public WebSocketHandlerMapping chenWebSocketHandlerMapping() { + var handlers = new LinkedHashMap(); + handlers.put("/ws/session", createRequestHandler(new SessionWebSocketHandler())); + handlers.put("/ws/console", createRequestHandler(new ConsoleWebSocketHandler())); + handlers.put("/ws/db-console", createRequestHandler(new DBConsoleWebsocketHandler())); + + var mapping = new WebSocketHandlerMapping(); + // 与 Spring 默认 WebSocket 映射一致,确保 WS 请求优先于普通 MVC 映射处理。 + mapping.setOrder(1); + mapping.setUrlMap(handlers); + return mapping; + } + + private WebSocketHttpRequestHandler createRequestHandler(WebSocketHandler webSocketHandler) { + var requestHandler = new WebSocketHttpRequestHandler(webSocketHandler); + // 仅使用 Chen 的动态校验,避免 WebSocketHandlerRegistry 追加第二个 Origin 拦截器。 + requestHandler.setHandshakeInterceptors(List.of(new ServletWebSocketHandshakeInterceptor())); + return requestHandler; } // 仅用于跨 Host 请求的精确白名单,格式为逗号分隔的 host 或 host:port。 @@ -146,6 +156,12 @@ public boolean beforeHandshake(ServerHttpRequest request, ServerHttpResponse res String origin = request.getHeaders().getOrigin(); InetSocketAddress requestHost = request.getHeaders().getHost(); + // 某些 Servlet 请求对象未填充 Host header,回退到容器解析到的主机和端口。 + if (requestHost == null && request instanceof ServletServerHttpRequest servletRequest) { + requestHost = new InetSocketAddress( + servletRequest.getServletRequest().getServerName(), + servletRequest.getServletRequest().getServerPort()); + } if (!checkOrigin(origin, requestHost, TRUSTED_DOMAINS)) { log.warn("Reject WebSocket handshake: untrusted or invalid origin"); @@ -161,6 +177,21 @@ public boolean beforeHandshake(ServerHttpRequest request, ServerHttpResponse res } var token = protocols.get(0); + var session = SessionManager.getSession(token); + var servletRequest = request instanceof ServletServerHttpRequest servletRequestWrapper + ? servletRequestWrapper.getServletRequest() : null; + var httpSession = servletRequest == null ? null : servletRequest.getSession(false); + var sessionBound = session != null && httpSession != null && Objects.equals( + session.getAttribute(SessionManager.WEB_SESSION_ID_ATTRIBUTE), httpSession.getId()); + // token 必须仍存活,且必须来自创建它的同一浏览器 HTTP session。 + if (!sessionBound) { + // 仅输出校验维度,不记录 token、Cookie 或 HTTP session ID 等敏感信息。 + log.warn("Reject WebSocket handshake: tokenExists={}, httpSessionExists={}, sessionBound={}", + session != null, httpSession != null, sessionBound); + response.setStatusCode(HttpStatus.UNAUTHORIZED); + return false; + } + log.info("Accept WebSocket handshake: HTTP session binding verified"); attributes.put("token", token); response.getHeaders().put("Sec-WebSocket-Protocol", protocols); return true; diff --git a/backend/web/src/main/java/org/jumpserver/chen/web/controller/AuthController.java b/backend/web/src/main/java/org/jumpserver/chen/web/controller/AuthController.java index 6e6691fe..34ebd361 100644 --- a/backend/web/src/main/java/org/jumpserver/chen/web/controller/AuthController.java +++ b/backend/web/src/main/java/org/jumpserver/chen/web/controller/AuthController.java @@ -32,6 +32,8 @@ public AuthResponse auth(HttpServletRequest request, @RequestBody AuthRequest au var lang = getLanguage(request); sess.setLocale(lang); + // Chen token 是 bearer token;记录所属浏览器 HTTP session,供后续 WS 握手校验。 + sess.setAttribute(SessionManager.WEB_SESSION_ID_ATTRIBUTE, request.getSession().getId()); var chenToken = SessionManager.registerSession(sess); return new AuthResponse(chenToken, lang.toLanguageTag()); diff --git a/backend/web/src/main/java/org/jumpserver/chen/web/controller/ConsoleController.java b/backend/web/src/main/java/org/jumpserver/chen/web/controller/ConsoleController.java index b936c015..d02a7368 100644 --- a/backend/web/src/main/java/org/jumpserver/chen/web/controller/ConsoleController.java +++ b/backend/web/src/main/java/org/jumpserver/chen/web/controller/ConsoleController.java @@ -13,6 +13,7 @@ import java.io.IOException; import java.nio.file.Files; +import java.nio.file.LinkOption; import java.time.LocalDateTime; import java.time.format.DateTimeFormatter; @@ -25,8 +26,18 @@ public ResponseEntity exportData(@PathVariable String fileKey) { if (!SessionManager.getCurrentSession().canDownload()) { throw new ChenException(MessageUtils.get("NoPermissionError")); } - var path = SessionManager.getCurrentSession().getTempPath(); - Resource resource = new FileSystemResource(path.resolve(fileKey).toFile()); + // export 文件名只能是单个文件名,先拒绝跨平台的路径分隔符和遍历片段。 + if (fileKey == null || fileKey.isBlank() + || fileKey.contains("/") || fileKey.contains("\\") || fileKey.contains("..")) { + throw new ChenException("Invalid export file"); + } + var basePath = SessionManager.getCurrentSession().getTempPath().toAbsolutePath().normalize(); + var filePath = basePath.resolve(fileKey).normalize(); + // normalize 后仍须落在会话目录内,且禁止末级符号链接指向目录外文件。 + if (!filePath.startsWith(basePath) || !Files.isRegularFile(filePath, LinkOption.NOFOLLOW_LINKS)) { + throw new ChenException("Invalid export file"); + } + Resource resource = new FileSystemResource(filePath.toFile()); var resp = ResponseEntity .ok() .contentType(MediaType.APPLICATION_OCTET_STREAM) diff --git a/backend/web/src/main/java/org/jumpserver/chen/web/service/ResourceService.java b/backend/web/src/main/java/org/jumpserver/chen/web/service/ResourceService.java index 79f2fbce..74a82a6b 100644 --- a/backend/web/src/main/java/org/jumpserver/chen/web/service/ResourceService.java +++ b/backend/web/src/main/java/org/jumpserver/chen/web/service/ResourceService.java @@ -5,11 +5,13 @@ import org.jumpserver.chen.framework.datasource.entity.action.Action; import org.jumpserver.chen.framework.datasource.entity.form.FormData; import org.jumpserver.chen.framework.session.SessionManager; +import org.jumpserver.chen.framework.utils.TreeUtils; import org.jumpserver.chen.web.exception.ChenException; import org.springframework.stereotype.Service; import java.sql.SQLException; import java.util.List; +import java.util.Objects; @Service @@ -32,12 +34,37 @@ public List getActions(TreeNode node) { public EventEmitter doAction(TreeNode node, String action) { try { var ds = SessionManager.getCurrentSession().getDatasource(); - return ds.doAction(node, action); + var resolvedNode = this.resolveActionNode(ds, node, action); + return ds.doAction(resolvedNode, action); } catch (Exception e) { + if (e instanceof ChenException) { + throw (ChenException) e; + } throw new ChenException(String.format("执行节点动作 %s 失败", node.getLabel()), e); } } + private TreeNode resolveActionNode(org.jumpserver.chen.framework.datasource.Datasource datasource, + TreeNode requestedNode, String action) throws SQLException { + if (requestedNode == null || requestedNode.getKey() == null || requestedNode.getKey().isBlank() + || action == null || action.isBlank()) { + throw new ChenException("Invalid resource action"); + } + + var root = datasource.getResourceBrowser().getTree(); + var resolvedNode = root == null ? null : TreeUtils.getNode(root, requestedNode.getKey()); + if (resolvedNode == null || !Objects.equals(resolvedNode.getType(), requestedNode.getType())) { + throw new ChenException("Invalid resource node"); + } + + var exposed = datasource.getActions(resolvedNode).stream() + .anyMatch(candidate -> Objects.equals(candidate.getKey(), action)); + if (!exposed) { + throw new ChenException("Invalid resource action"); + } + return resolvedNode; + } + public EventEmitter submitResourceForm(FormData form) throws SQLException { var ds = SessionManager.getCurrentSession().getDatasource(); return ds.handleForm(form); diff --git a/backend/web/src/main/java/org/jumpserver/chen/web/service/impl/JmsSessionService.java b/backend/web/src/main/java/org/jumpserver/chen/web/service/impl/JmsSessionService.java index 15123c36..306f1552 100644 --- a/backend/web/src/main/java/org/jumpserver/chen/web/service/impl/JmsSessionService.java +++ b/backend/web/src/main/java/org/jumpserver/chen/web/service/impl/JmsSessionService.java @@ -17,6 +17,7 @@ import java.net.InetAddress; import java.net.UnknownHostException; import java.time.Instant; +import java.util.Map; @Service @Slf4j @@ -129,14 +130,7 @@ private Datasource createDatasource(ServiceOuterClass.TokenResponse tokenResp) { dbConnectInfo.setDb(tokenResp.getData().getAsset().getSpecific().getDbName()); var platformSettings = tokenResp.getData().getPlatform().getProtocols(0).getSettingsMap(); -// - if (platformSettings.containsKey("sysdba") && platformSettings.get("sysdba").equals("true")) { - dbConnectInfo.getOptions().put("internal_logon", "sysdba"); - } - - if (platformSettings.containsKey("version")) { - dbConnectInfo.getOptions().put("version", platformSettings.get("version")); - } + applyPlatformSettings(dbConnectInfo, platformSettings); var asset = tokenResp.getData().getAsset(); @@ -151,6 +145,21 @@ private Datasource createDatasource(ServiceOuterClass.TokenResponse tokenResp) { return DatasourceFactory.fromConnectInfo(dbConnectInfo); } + static void applyPlatformSettings(DBConnectInfo dbConnectInfo, Map platformSettings) { + if (platformSettings.containsKey("sysdba") && platformSettings.get("sysdba").equals("true")) { + dbConnectInfo.getOptions().put("internal_logon", "sysdba"); + } + + if (platformSettings.containsKey("version")) { + dbConnectInfo.getOptions().put("version", platformSettings.get("version")); + } + + if ("sqlserver".equals(dbConnectInfo.getDbType())) { + var encrypt = platformSettings.getOrDefault("encrypt", "true"); + dbConnectInfo.getOptions().put("encrypt", Boolean.parseBoolean(encrypt)); + } + } + private Common.Session createJMSSession(ServiceOuterClass.TokenResponse tokenResp, String remoteAddr) { var jmsSession = Common.Session.newBuilder() .setUserId(tokenResp.getData().getUser().getId()) diff --git a/frontend/src/components/Main/Explore/Dialog/FormTemplate.vue b/frontend/src/components/Main/Explore/Dialog/FormTemplate.vue index 9d9ea97a..ebd2f98a 100644 --- a/frontend/src/components/Main/Explore/Dialog/FormTemplate.vue +++ b/frontend/src/components/Main/Explore/Dialog/FormTemplate.vue @@ -52,28 +52,7 @@ import { compileSQL } from '@/utils/sql' import 'codemirror/theme/3024-night.css' import 'codemirror/mode/sql/sql.js' import store from '@/store' -import { format } from 'sql-formatter' - -const formatterMap = { - null: 'sql', - 'mariadb': 'mariadb', - 'mysql': 'mysql', - 'postgresql': 'postgresql', - 'oracle': 'plsql', - 'db2': 'db2', - 'dameng': 'dameng' -} - -const modeMap = { - null: 'text/x-sql', - 'mariadb': 'text/x-mariadb', - 'mysql': 'text/x-mysql', - 'postgresql': 'text/x-pgsql', - 'oracle': 'text/x-plsql', - 'sqlserver': 'text/x-mssql', - 'db2': 'text/x-sql', - 'dameng': 'text/x-sql' -} +import { formatSqlForEditor, getEditorMode } from '@/utils/sqlEditorSupport' export default { props: { @@ -102,7 +81,7 @@ export default { lineNumbers: true, line: true, readOnly: true, - mode: modeMap[store.getters.profile?.dbType], + mode: getEditorMode(store.getters.profile?.dbType), cursorBlinkRate: -1 } } @@ -132,8 +111,7 @@ export default { const sql = compileSQL( this.formMeta.sqlTemplate, this.sqlParams ) - const lang = formatterMap[store.getters.profile?.dbType] - return format(sql, { language: lang }) + return formatSqlForEditor(sql, store.getters.profile?.dbType) }, set() {} } diff --git a/frontend/src/components/Main/Explore/QueryConsole/CodeEditor.vue b/frontend/src/components/Main/Explore/QueryConsole/CodeEditor.vue index 76adcaa2..e09475b8 100644 --- a/frontend/src/components/Main/Explore/QueryConsole/CodeEditor.vue +++ b/frontend/src/components/Main/Explore/QueryConsole/CodeEditor.vue @@ -46,10 +46,10 @@