From 7705354142a25f6461be380e9a68ffbad41df543 Mon Sep 17 00:00:00 2001 From: HY-love-sleep <583699747@qq.com> Date: Sun, 20 Sep 2026 17:47:27 +0800 Subject: [PATCH] fix: AI token limiter null handling, rule isolation, and gzip streaming (#6513, #6514, #6515) 1. #6513 null rule fields crash request processing Fill defaults for tokenLimit, timeWindowSeconds, aiTokenLimitType and keyName in AiTokenLimiterPluginHandler.handlerRule(), reusing AiTokenLimiterHandle.newDefaultInstance() so the values match the defaults the admin applies to a new rule. The crash came from isAllowed() (unboxing a null tokenLimit) and from recordTokensUsage() (Duration.ofSeconds(null)). 2. #6514 Redis counter key ignores selector and rule scope The key becomes PREFIX + CacheKeyUtils.INST.getKey(rule) + ":" + resolverValue, reusing the same rule key as the handle cache, so counters are scoped per selector and rule. Note: the key format change means existing counters expire naturally via TTL. 3. #6515 gzip responses spanning multiple DataBuffers undercount tokens The previous implementation located the gzip header by hand and assumed it was fully contained in the first DataBuffer; when the header straddled a buffer boundary the computed offset was wrong and the payload was never parsed. Add a package-visible GzipStreamDecoder: GzipHeaderState parses the gzip header across buffers (FEXTRA/FNAME/FCOMMENT/FHCRC) and an incremental Inflater decompresses each buffer, keeping memory bounded - in the same spirit as #7124. Known limitations: multi-member gzip streams are not supported (decompression stops after the first deflate stream), and a header larger than 266 bytes abandons decompression with a WARN. Tests: 7 cases covering a header spanning buffers, a header contained in a single buffer, a regression case whose first chunk is far larger than the header buffer, and multi-chunk end-to-end decompression asserted by exact string comparison. --- .../token/limiter/AiTokenLimiterPlugin.java | 119 ++++------- .../ai/token/limiter/GzipStreamDecoder.java | 191 ++++++++++++++++++ .../handler/AiTokenLimiterPluginHandler.java | 14 ++ .../limiter/AiTokenLimiterPluginTest.java | 144 +++++++++++++ 4 files changed, 385 insertions(+), 83 deletions(-) create mode 100644 shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/main/java/org/apache/shenyu/plugin/ai/token/limiter/GzipStreamDecoder.java diff --git a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/main/java/org/apache/shenyu/plugin/ai/token/limiter/AiTokenLimiterPlugin.java b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/main/java/org/apache/shenyu/plugin/ai/token/limiter/AiTokenLimiterPlugin.java index 74ae71078386..0094a1431373 100644 --- a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/main/java/org/apache/shenyu/plugin/ai/token/limiter/AiTokenLimiterPlugin.java +++ b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/main/java/org/apache/shenyu/plugin/ai/token/limiter/AiTokenLimiterPlugin.java @@ -51,7 +51,6 @@ import reactor.core.publisher.Mono; import reactor.util.annotation.NonNull; -import java.io.ByteArrayOutputStream; import java.nio.ByteBuffer; import java.nio.charset.StandardCharsets; import java.time.Duration; @@ -61,8 +60,6 @@ import java.util.function.Consumer; import java.util.regex.Matcher; import java.util.regex.Pattern; -import java.util.zip.DataFormatException; -import java.util.zip.Inflater; /** * Shenyu ai token limiter plugin. @@ -88,13 +85,13 @@ protected Mono doExecute(final ServerWebExchange exchange, final ShenyuPlu ReactiveRedisTemplate reactiveRedisTemplate = AiTokenLimiterPluginHandler.REDIS_CACHED_HANDLE.get().obtainHandle(PluginEnum.AI_TOKEN_LIMITER.getName()); Assert.notNull(reactiveRedisTemplate, "reactiveRedisTemplate is null"); - // generate redis key + // generate redis key - include rule id to scope counters per rule String tokenLimitType = aiTokenLimiterHandle.getAiTokenLimitType(); String keyName = aiTokenLimiterHandle.getKeyName(); Long tokenLimit = aiTokenLimiterHandle.getTokenLimit(); Long timeWindowSeconds = aiTokenLimiterHandle.getTimeWindowSeconds(); - String cacheKey = REDIS_KEY_PREFIX + getCacheKey(exchange, tokenLimitType, keyName); + String cacheKey = REDIS_KEY_PREFIX + CacheKeyUtils.INST.getKey(rule) + ":" + getCacheKey(exchange, tokenLimitType, keyName); final AiStatisticServerHttpResponse loggingServerHttpResponse = new AiStatisticServerHttpResponse(exchange, exchange.getResponse(), tokens -> recordTokensUsage(reactiveRedisTemplate, @@ -229,9 +226,7 @@ private Flux appendResponse(final Publisherfrom(body) .doOnNext(buffer -> { @@ -241,59 +236,18 @@ private Flux appendResponse(final Publisher 0) { - baos.write(outBuf, 0, cnt); - } - } catch (DataFormatException ex) { - LOG.error("Inflater decompression failed", ex); - } - processedBytes = baos.toByteArray(); - } else { - processedBytes = inBytes; + byte[] processedBytes = isGzip ? decoder.decode(inBytes) : inBytes; + if (processedBytes.length > 0) { + processChunk(processedBytes, writer); } - String chunk = new String(processedBytes, StandardCharsets.UTF_8); - for (String line : chunk.split("\\r?\\n")) { - if (!line.startsWith("data:")) { - continue; - } - String payload = line.substring("data:".length()).trim(); - if (payload.isEmpty() || "[DONE]".equals(payload)) { - continue; - } - if (!payload.startsWith("{")) { - continue; - } - try { - JsonNode node = MAPPER.readTree(payload); - JsonNode usage = node.get(Constants.USAGE); - if (Objects.nonNull(usage) && usage.has(Constants.COMPLETION_TOKENS)) { - long c = usage.get(Constants.COMPLETION_TOKENS).asLong(); - tokensRecorder.accept(c); - streamingUsageRecorded.set(true); - } - } catch (Exception e) { - LOG.error("Failed to parse AI response JSON payload", e); - } - } - writer.write(ByteBuffer.wrap(processedBytes)); }); } catch (Exception e) { LOG.error("read dataBuffer error", e); } }) .doFinally(signal -> { - if (Objects.nonNull(inflater)) { - inflater.end(); + if (Objects.nonNull(decoder)) { + decoder.close(); } if (!streamingUsageRecorded.get()) { String sse = writer.output(); @@ -303,6 +257,34 @@ private Flux appendResponse(final Publisher 0) { + baos.write(decompressBuffer, 0, cnt); + } + } catch (DataFormatException ex) { + LOG.error("Inflater decompression failed", ex); + abandoned = true; + return new byte[0]; + } + return baos.toByteArray(); + } + + void close() { + inflater.end(); + } + + /** + * Track gzip header parsing state across buffers. + */ + static class GzipHeaderState { + + private static final int MAX_HEADER_SIZE = 10 + 256; + + private final byte[] accumulatedHeader = new byte[MAX_HEADER_SIZE]; + + private int accumulatedLength; + + private boolean complete; + + private boolean capacityExceeded; + + boolean isComplete() { + return complete; + } + + boolean isCapacityExceeded() { + return capacityExceeded; + } + + /** + * Process gzip header bytes, potentially spanning multiple buffers. + * + * @param inBytes input bytes + * @return offset where compressed data starts (0 if header is still incomplete) + */ + int process(final byte[] inBytes) { + if (complete || capacityExceeded) { + return 0; + } + + final int prev = accumulatedLength; + final int toCopy = Math.min(inBytes.length, accumulatedHeader.length - prev); + System.arraycopy(inBytes, 0, accumulatedHeader, prev, toCopy); + accumulatedLength += toCopy; + + if (accumulatedLength < 10) { + return 0; + } + + try { + int pos = 10; + int flg = accumulatedHeader[3] & 0xFF; + + if ((flg & 0x04) != 0) { + if (accumulatedLength < pos + 2) { + return headerIncomplete(); + } + int xlen = (accumulatedHeader[pos] & 0xFF) | ((accumulatedHeader[pos + 1] & 0xFF) << 8); + pos += 2 + xlen; + if (accumulatedLength < pos) { + return headerIncomplete(); + } + } + + if ((flg & 0x08) != 0) { + while (pos < accumulatedLength && accumulatedHeader[pos] != 0) { + pos++; + } + if (pos >= accumulatedLength) { + return headerIncomplete(); + } + pos++; + } + + if ((flg & 0x10) != 0) { + while (pos < accumulatedLength && accumulatedHeader[pos] != 0) { + pos++; + } + if (pos >= accumulatedLength) { + return headerIncomplete(); + } + pos++; + } + + if ((flg & 0x02) != 0) { + if (accumulatedLength < pos + 2) { + return headerIncomplete(); + } + pos += 2; + } + + complete = true; + return pos - prev; + + } catch (ArrayIndexOutOfBoundsException e) { + // Defensive: the bounds checks above should make this unreachable. + // Abandon decompression instead of letting the error propagate into the + // reactive pipeline, where it would abort the response for the client. + capacityExceeded = true; + LOG.warn("Unexpected gzip header parse error, decompression abandoned", e); + return 0; + } + } + + private int headerIncomplete() { + if (accumulatedLength >= accumulatedHeader.length) { + capacityExceeded = true; + LOG.warn("Gzip header exceeds maximum size of {} bytes, decompression abandoned. " + + "This may occur with long FNAME or FCOMMENT fields.", MAX_HEADER_SIZE); + } + return 0; + } + } +} diff --git a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/main/java/org/apache/shenyu/plugin/ai/token/limiter/handler/AiTokenLimiterPluginHandler.java b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/main/java/org/apache/shenyu/plugin/ai/token/limiter/handler/AiTokenLimiterPluginHandler.java index 5408872e3ba6..b97738471e03 100644 --- a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/main/java/org/apache/shenyu/plugin/ai/token/limiter/handler/AiTokenLimiterPluginHandler.java +++ b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/main/java/org/apache/shenyu/plugin/ai/token/limiter/handler/AiTokenLimiterPluginHandler.java @@ -88,6 +88,20 @@ public void removeSelector(final SelectorData selectorData) { public void handlerRule(final RuleData ruleData) { Optional.ofNullable(ruleData.getHandle()).ifPresent(s -> { final AiTokenLimiterHandle rateLimiterHandle = GsonUtils.getInstance().fromJson(s, AiTokenLimiterHandle.class); + // Fill defaults for null fields to prevent NPE + AiTokenLimiterHandle defaultHandle = AiTokenLimiterHandle.newDefaultInstance(); + if (Objects.isNull(rateLimiterHandle.getTokenLimit())) { + rateLimiterHandle.setTokenLimit(defaultHandle.getTokenLimit()); + } + if (Objects.isNull(rateLimiterHandle.getTimeWindowSeconds())) { + rateLimiterHandle.setTimeWindowSeconds(defaultHandle.getTimeWindowSeconds()); + } + if (Objects.isNull(rateLimiterHandle.getAiTokenLimitType())) { + rateLimiterHandle.setAiTokenLimitType(defaultHandle.getAiTokenLimitType()); + } + if (Objects.isNull(rateLimiterHandle.getKeyName())) { + rateLimiterHandle.setKeyName(defaultHandle.getKeyName()); + } CACHED_HANDLE.get().cachedHandle(CacheKeyUtils.INST.getKey(ruleData), rateLimiterHandle); }); } diff --git a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/test/java/org/apache/shenyu/plugin/ai/token/limiter/AiTokenLimiterPluginTest.java b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/test/java/org/apache/shenyu/plugin/ai/token/limiter/AiTokenLimiterPluginTest.java index d929206d80eb..3142a6049796 100644 --- a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/test/java/org/apache/shenyu/plugin/ai/token/limiter/AiTokenLimiterPluginTest.java +++ b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-token-limiter/src/test/java/org/apache/shenyu/plugin/ai/token/limiter/AiTokenLimiterPluginTest.java @@ -17,13 +17,21 @@ package org.apache.shenyu.plugin.ai.token.limiter; +import org.apache.shenyu.common.dto.RuleData; +import org.apache.shenyu.common.dto.convert.rule.AiTokenLimiterHandle; +import org.apache.shenyu.plugin.ai.token.limiter.handler.AiTokenLimiterPluginHandler; +import org.apache.shenyu.plugin.base.utils.CacheKeyUtils; import org.junit.jupiter.api.Test; +import java.io.ByteArrayOutputStream; +import java.io.IOException; import java.nio.ByteBuffer; import java.nio.charset.StandardCharsets; +import java.util.zip.GZIPOutputStream; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; import static org.junit.jupiter.api.Assertions.assertTrue; /** @@ -51,4 +59,140 @@ void testBodyWriterKeepsTailOfLargeChunk() { assertEquals("23456789", writer.output()); } + + @Test + void testHandlerRuleWithNullFieldsFillsDefaults() { + // Test for issue #6513: null fields should be filled with defaults + AiTokenLimiterPluginHandler handler = new AiTokenLimiterPluginHandler(); + + RuleData ruleData = RuleData.builder() + .id("test-rule-id") + .selectorId("test-selector-id") + .name("test-rule") + .handle("{\"aiTokenLimitType\":\"uri\",\"keyName\":\"default\"}") + .build(); + + handler.handlerRule(ruleData); + + AiTokenLimiterHandle cached = AiTokenLimiterPluginHandler.CACHED_HANDLE.get() + .obtainHandle(CacheKeyUtils.INST.getKey(ruleData)); + + assertNotNull(cached); + assertEquals("uri", cached.getAiTokenLimitType()); + assertEquals("default", cached.getKeyName()); + // These should be filled with defaults + assertNotNull(cached.getTokenLimit()); + assertNotNull(cached.getTimeWindowSeconds()); + assertEquals(Long.valueOf(100L), cached.getTokenLimit()); + assertEquals(Long.valueOf(60L), cached.getTimeWindowSeconds()); + } + + @Test + void testGzipDecoderWithHeaderSpanningBuffers() throws IOException { + // Test for issue #6515: gzip header spanning multiple buffers + String sseContent = "data: {\"id\":\"test\",\"usage\":{\"completion_tokens\":50}}\n\n"; + byte[] compressed = compressGzip(sseContent); + + // Split at byte 8 - header boundary + byte[] chunk1 = new byte[8]; + byte[] chunk2 = new byte[compressed.length - 8]; + System.arraycopy(compressed, 0, chunk1, 0, 8); + System.arraycopy(compressed, 8, chunk2, 0, chunk2.length); + + GzipStreamDecoder decoder = new GzipStreamDecoder(); + byte[] result1 = decoder.decode(chunk1); + // Header incomplete + assertEquals(0, result1.length); + + byte[] result2 = decoder.decode(chunk2); + // Should decompress now + assertTrue(result2.length > 0); + assertEquals(sseContent, new String(result2, StandardCharsets.UTF_8)); + decoder.close(); + } + + @Test + void testGzipDecoderWithCompleteHeaderInFirstBuffer() throws IOException { + String sseContent = "data: {\"id\":\"test\",\"usage\":{\"completion_tokens\":50}}\n\n"; + byte[] compressed = compressGzip(sseContent); + + GzipStreamDecoder decoder = new GzipStreamDecoder(); + byte[] result = decoder.decode(compressed); + assertEquals(sseContent, new String(result, StandardCharsets.UTF_8)); + decoder.close(); + } + + @Test + void testGzipDecoderWithLargeFirstChunk() throws IOException { + // Regression: first chunk far larger than header buffer limit (266) + // historically would incorrectly abandon entire response + StringBuilder sb = new StringBuilder(); + for (int i = 0; i < 200; i++) { + sb.append("data: {\"id\":\"chatcmpl-").append(i) + .append("\",\"choices\":[{\"delta\":{\"content\":\"token-").append(i * 7919).append("\"}}]}\n\n"); + } + sb.append("data: {\"usage\":{\"completion_tokens\":75}}\n\n"); + String sseContent = sb.toString(); + + byte[] compressed = compressGzip(sseContent); + assertTrue(compressed.length > 266, "Compressed data must exceed header buffer size"); + + GzipStreamDecoder decoder = new GzipStreamDecoder(); + byte[] result = decoder.decode(compressed); + decoder.close(); + + assertEquals(sseContent, new String(result, StandardCharsets.UTF_8)); + } + + @Test + void testEndToEndGzipDecompressionAcrossMultipleChunks() throws IOException { + // End-to-end test: verify full decompression pipeline with 3 chunks + String sseContent = buildSseContentWithTokens(75); + byte[] compressed = compressGzip(sseContent); + byte[][] chunks = splitIntoThreeChunks(compressed); + + GzipStreamDecoder decoder = new GzipStreamDecoder(); + StringBuilder decompressed = new StringBuilder(); + + for (byte[] chunk : chunks) { + byte[] result = decoder.decode(chunk); + if (result.length > 0) { + decompressed.append(new String(result, StandardCharsets.UTF_8)); + } + } + decoder.close(); + + // Verify decompression succeeded + assertEquals(sseContent, decompressed.toString()); + assertTrue(decompressed.toString().contains("completion_tokens\":75")); + } + + private String buildSseContentWithTokens(final int tokens) { + return "data: {\"id\":\"chatcmpl-1\",\"choices\":[{\"delta\":{\"content\":\"Hello\"}}]}\n\n" + + "data: {\"id\":\"chatcmpl-1\",\"choices\":[{\"delta\":{\"content\":\" World\"}}]}\n\n" + + "data: {\"id\":\"chatcmpl-1\",\"choices\":[{\"delta\":{}}]," + + "\"usage\":{\"completion_tokens\":" + tokens + ",\"prompt_tokens\":10,\"total_tokens\":" + (tokens + 10) + "}}\n\n" + + "data: [DONE]\n\n"; + } + + private byte[] compressGzip(final String content) throws IOException { + ByteArrayOutputStream compressedStream = new ByteArrayOutputStream(); + try (GZIPOutputStream gzipOutputStream = new GZIPOutputStream(compressedStream)) { + gzipOutputStream.write(content.getBytes(StandardCharsets.UTF_8)); + } + return compressedStream.toByteArray(); + } + + private byte[][] splitIntoThreeChunks(final byte[] data) { + byte[] chunk1 = new byte[8]; + int chunk2Size = (data.length - 8) / 2; + byte[] chunk2 = new byte[chunk2Size]; + byte[] chunk3 = new byte[data.length - 8 - chunk2Size]; + + System.arraycopy(data, 0, chunk1, 0, 8); + System.arraycopy(data, 8, chunk2, 0, chunk2Size); + System.arraycopy(data, 8 + chunk2Size, chunk3, 0, chunk3.length); + + return new byte[][]{chunk1, chunk2, chunk3}; + } }