本文目录导读:

我将为您提供一个Java异步线程上下文的完整案例,涵盖常见的场景和解决方案。
基础案例:ThreadLocal在异步场景的问题
import java.util.concurrent.*;
import java.util.HashMap;
import java.util.Map;
/**
* 演示ThreadLocal在异步场景中的丢失问题
*/
public class AsyncThreadLocalDemo {
// 模拟用户上下文
private static ThreadLocal<String> userContext = new ThreadLocal<>();
private static ThreadLocal<Map<String, Object>> requestContext = new ThreadLocal<>();
public static void main(String[] args) throws Exception {
ExecutorService executor = Executors.newFixedThreadPool(5);
// 在主线程设置上下文
userContext.set("User-001");
Map<String, Object> attrs = new HashMap<>();
attrs.put("requestId", "REQ-123");
attrs.put("traceId", "TRACE-456");
requestContext.set(attrs);
System.out.println("主线程设置: " + userContext.get());
System.out.println("主线程请求上下文: " + requestContext.get());
// 使用submit提交任务
Future<String> future = executor.submit(() -> {
// 子线程中无法获取主线程的ThreadLocal
return "子线程获取: " + userContext.get();
});
System.out.println(future.get());
// 使用execute提交任务
executor.execute(() -> {
System.out.println("execute方式获取: " + userContext.get());
});
executor.shutdown();
}
}
解决方案:自定义ThreadLocal包装器
import java.util.HashMap;
import java.util.Map;
import java.util.concurrent.*;
/**
* 可传递的ThreadLocal包装器
*/
public class TransmittableThreadLocal<T> {
// 存储当前线程的变量
private static final ThreadLocal<Map<String, Object>> THREAD_LOCAL = new ThreadLocal<>();
// 值提供者
private final String key;
private final T defaultValue;
public TransmittableThreadLocal(String key, T defaultValue) {
this.key = key;
this.defaultValue = defaultValue;
}
public void set(T value) {
Map<String, Object> map = getContextMap();
map.put(key, value);
}
@SuppressWarnings("unchecked")
public T get() {
Map<String, Object> map = getContextMap();
T value = (T) map.get(key);
return value != null ? value : defaultValue;
}
public void remove() {
Map<String, Object> map = getContextMap();
map.remove(key);
}
private Map<String, Object> getContextMap() {
Map<String, Object> map = THREAD_LOCAL.get();
if (map == null) {
map = new HashMap<>();
THREAD_LOCAL.set(map);
}
return map;
}
/**
* 获取整个上下文
*/
public static Map<String, Object> capture() {
Map<String, Object> map = THREAD_LOCAL.get();
return map != null ? new ConcurrentHashMap<>(map) : null;
}
/**
* 设置整个上下文
*/
public static void replay(Map<String, Object> context) {
if (context != null) {
THREAD_LOCAL.set(new ConcurrentHashMap<>(context));
}
}
/**
* 清除上下文
*/
public static void clear() {
THREAD_LOCAL.remove();
}
}
完整的异步上下文传递方案
import java.util.*;
import java.util.concurrent.*;
import java.util.function.Supplier;
/**
* 完整的异步上下文传递解决方案
*/
public class AsyncContextDemo {
private final ExecutorService executor;
private final TransmittableThreadLocal<String> userId;
private final TransmittableThreadLocal<String> traceId;
private final TransmittableThreadLocal<Map<String, Object>> attributes;
public AsyncContextDemo() {
this.executor = Executors.newFixedThreadPool(10);
this.userId = new TransmittableThreadLocal<>("userId", "unknown");
this.traceId = new TransmittableThreadLocal<>("traceId", "no-trace");
this.attributes = new TransmittableThreadLocal<>("attrs", new HashMap<>());
}
/**
* 包装Runnable,传递上下文
*/
public Runnable wrapRunnable(Runnable task) {
// 捕获当前上下文
Map<String, Object> context = TransmittableThreadLocal.capture();
return () -> {
try {
// 在子线程中重放上下文
TransmittableThreadLocal.replay(context);
task.run();
} finally {
// 清理子线程上下文
TransmittableThreadLocal.clear();
}
};
}
/**
* 包装Callable,传递上下文
*/
public <T> Callable<T> wrapCallable(Callable<T> task) {
Map<String, Object> context = TransmittableThreadLocal.capture();
return () -> {
try {
TransmittableThreadLocal.replay(context);
return task.call();
} finally {
TransmittableThreadLocal.clear();
}
};
}
/**
* 包装Supplier,传递上下文
*/
public <T> Supplier<T> wrapSupplier(Supplier<T> supplier) {
Map<String, Object> context = TransmittableThreadLocal.capture();
return () -> {
try {
TransmittableThreadLocal.replay(context);
return supplier.get();
} finally {
TransmittableThreadLocal.clear();
}
};
}
/**
* 提交带上下文的任务
*/
public Future<?> submit(Runnable task) {
return executor.submit(wrapRunnable(task));
}
public <T> Future<T> submit(Callable<T> task) {
return executor.submit(wrapCallable(task));
}
public void execute(Runnable task) {
executor.execute(wrapRunnable(task));
}
/**
* 测试示例
*/
public void testDemo() throws Exception {
// 设置上下文
userId.set("User-100");
traceId.set("TRACE-123456");
attributes.get().put("requestId", "REQ-ABCDEF");
System.out.println("主线程开始 - User: " + userId.get()
+ ", Trace: " + traceId.get());
// 示例1:使用submit执行
Future<String> future = submit(() -> {
// 模拟业务操作
Thread.sleep(100);
return "任务1 - User: " + userId.get()
+ ", Trace: " + traceId.get()
+ ", Attr: " + attributes.get().get("requestId");
});
System.out.println("子线程结果: " + future.get());
// 示例2:使用execute执行
execute(() -> {
System.out.println("任务2 - User: " + userId.get());
});
// 示例3:使用CompletableFuture
CompletableFuture<String> cf = CompletableFuture.supplyAsync(
wrapSupplier(() -> "CompletableFuture - User: " + userId.get()),
executor
);
System.out.println("CF结果: " + cf.get());
// 示例4:异常处理
Future<String> errorFuture = submit(() -> {
if (true) throw new RuntimeException("模拟异常");
return "永远执行不到";
});
try {
errorFuture.get();
} catch (Exception e) {
System.out.println("捕获异常: " + e.getCause().getMessage());
}
executor.shutdown();
}
public static void main(String[] args) throws Exception {
new AsyncContextDemo().testDemo();
}
}
使用现成库(TransmittableThreadLocal)
import com.alibaba.ttl.TransmittableThreadLocal;
import com.alibaba.ttl.threadpool.TtlExecutors;
import java.util.concurrent.*;
/**
* 使用阿里巴巴TTL库
*/
public class TTLDemo {
// 使用TTL的ThreadLocal
private static TransmittableThreadLocal<String> context =
new TransmittableThreadLocal<>();
public static void main(String[] args) throws Exception {
// 创建线程池
ExecutorService executor = Executors.newFixedThreadPool(5);
// 使用TtlExecutors包装线程池
ExecutorService ttlExecutor = TtlExecutors.getTtlExecutorService(executor);
// 主线程设置上下文
context.set("TTL-Context-Value");
System.out.println("主线程: " + context.get());
// 提交任务
ttlExecutor.execute(() -> {
System.out.println("子线程获取: " + context.get());
});
// 使用FutureTask
Future<String> future = ttlExecutor.submit(() -> {
return "子线程返回值: " + context.get();
});
System.out.println(future.get());
// 异步任务完成后再修改上下文
Thread.sleep(1000);
context.set("Modified-Value");
ttlExecutor.execute(() -> {
System.out.println("修改后的传递: " + context.get());
});
ttlExecutor.shutdown();
}
}
实际业务应用案例
import java.util.*;
import java.util.concurrent.*;
/**
* 业务场景:分布式追踪系统上下文传递
*/
public class BusinessContextDemo {
// 业务上下文
private static final ThreadLocal<TraceContext> TRACE_CONTEXT = new ThreadLocal<>();
// 业务上下文类
static class TraceContext {
private final String traceId;
private final String spanId;
private final String userId;
private final Map<String, String> tags = new HashMap<>();
public TraceContext(String traceId, String spanId, String userId) {
this.traceId = traceId;
this.spanId = spanId;
this.userId = userId;
}
@Override
public String toString() {
return String.format("TraceContext{traceId='%s', spanId='%s', userId='%s', tags=%s}",
traceId, spanId, userId, tags);
}
}
/**
* 自定义线程池工厂,自动传递上下文
*/
static class ContextAwareThreadFactory implements ThreadFactory {
private final ThreadFactory delegate;
public ContextAwareThreadFactory(ThreadFactory delegate) {
this.delegate = delegate;
}
@Override
public Thread newThread(Runnable r) {
Thread thread = delegate.newThread(r);
// 捕获当前请求上下文
TraceContext context = TRACE_CONTEXT.get();
if (context != null) {
thread.setUncaughtExceptionHandler((t, e) -> {
System.err.println("线程异常: " + e.getMessage());
});
}
return thread;
}
}
/**
* 异步任务管理器
*/
static class AsyncTaskManager {
private final ExecutorService executor;
public AsyncTaskManager() {
this.executor = createExecutor();
}
private ExecutorService createExecutor() {
ThreadFactory factory = new ContextAwareThreadFactory(
Executors.defaultThreadFactory()
);
return Executors.newFixedThreadPool(10, factory);
}
/**
* 提交带上下文的异步任务
*/
public <T> CompletableFuture<T> submitAsync(Callable<T> task) {
TraceContext context = TRACE_CONTEXT.get();
return CompletableFuture.supplyAsync(() -> {
try {
// 在子线程恢复上下文
TraceContext oldContext = TRACE_CONTEXT.get();
TRACE_CONTEXT.set(context);
try {
return task.call();
} finally {
if (oldContext != null) {
TRACE_CONTEXT.set(oldContext);
} else {
TRACE_CONTEXT.remove();
}
}
} catch (Exception e) {
throw new RuntimeException(e);
}
}, executor);
}
public void shutdown() {
executor.shutdown();
}
}
public static void main(String[] args) throws Exception {
// 创建异步任务管理器
AsyncTaskManager manager = new AsyncTaskManager();
// 模拟一个请求
String traceId = UUID.randomUUID().toString();
String spanId = UUID.randomUUID().toString().substring(0, 8);
String userId = "User-" + System.currentTimeMillis() % 1000;
TraceContext context = new TraceContext(traceId, spanId, userId);
context.tags.put("client", "web-app");
context.tags.put("env", "production");
// 设置请求上下文
TRACE_CONTEXT.set(context);
System.out.println("主线程开始处理请求: " + context);
// 执行多个异步任务
CompletableFuture<String> task1 = manager.submitAsync(() -> {
Thread.sleep(100);
TraceContext ctx = TRACE_CONTEXT.get();
return "任务1完成 - " + ctx;
});
CompletableFuture<String> task2 = manager.submitAsync(() -> {
Thread.sleep(200);
TraceContext ctx = TRACE_CONTEXT.get();
return "任务2完成 - " + ctx;
});
CompletableFuture<String> task3 = manager.submitAsync(() -> {
Thread.sleep(150);
TraceContext ctx = TRACE_CONTEXT.get();
return "任务3完成 - " + ctx;
});
// 等待所有任务完成
CompletableFuture.allOf(task1, task2, task3).join();
System.out.println("任务1: " + task1.get());
System.out.println("任务2: " + task2.get());
System.out.println("任务3: " + task3.get());
// 主线程上下文仍然保持
System.out.println("主线程上下文: " + TRACE_CONTEXT.get());
// 清理
TRACE_CONTEXT.remove();
manager.shutdown();
}
}
-
问题:
ThreadLocal默认不会在线程间传递,子线程无法继承父线程的上下文。 -
解决方案:
- 手动通过构造函数传递
- 使用包装器模式包装Runnable/Callable
- 使用阿里巴巴的
TransmittableThreadLocal - 自定义线程池工厂
-
最佳实践:
- 在异步任务入口捕获上下文
- 在子线程中恢复上下文
- 确保finally中清理上下文
- 注意线程池复用导致的上下文混淆
-
性能考虑:
- 避免传递大对象
- 使用合适的并发容器
- 考虑内存泄漏风险
这些方案可以有效解决Java异步编程中的上下文传递问题,适用于分布式追踪、用户认证、请求日志等场景。