Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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.
Expand All @@ -88,13 +85,13 @@ protected Mono<Void> 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,
Expand Down Expand Up @@ -229,9 +226,7 @@ private Flux<? extends DataBuffer> appendResponse(final Publisher<? extends Data
&& headers.getFirst(Constants.CONTENT_ENCODING)
.contains(Constants.HTTP_ACCEPT_ENCODING_GZIP);

final Inflater inflater = isGzip ? new Inflater(true) : null;
final byte[] outBuf = new byte[4096];
final AtomicBoolean headerSkipped = new AtomicBoolean(!isGzip);
final GzipStreamDecoder decoder = isGzip ? new GzipStreamDecoder() : null;

return Flux.<DataBuffer>from(body)
.doOnNext(buffer -> {
Expand All @@ -241,59 +236,18 @@ private Flux<? extends DataBuffer> appendResponse(final Publisher<? extends Data
byte[] inBytes = new byte[ro.remaining()];
ro.get(inBytes);

byte[] processedBytes;
if (isGzip) {
int offset = 0;
if (headerSkipped.compareAndSet(false, true)) {
offset = skipGzipHeader(inBytes);
}
inflater.setInput(inBytes, offset, inBytes.length - offset);
ByteArrayOutputStream baos = new ByteArrayOutputStream();
try {
int cnt;
while ((cnt = inflater.inflate(outBuf)) > 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();
Expand All @@ -303,6 +257,34 @@ private Flux<? extends DataBuffer> appendResponse(final Publisher<? extends Data
});
}

private void processChunk(final byte[] processedBytes, final BodyWriter 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));
}

private long extractUsageTokensFromSse(final String sse) {
Matcher m = COMPLETION_TOKENS_PATTERN.matcher(sse);
long last = 0L;
Expand All @@ -312,35 +294,6 @@ private long extractUsageTokensFromSse(final String sse) {
return last;
}

private int skipGzipHeader(final byte[] b) {
int pos = 10;
int flg = b[3] & 0xFF;

if ((flg & 0x04) != 0) {
int xlen = (b[pos] & 0xFF) | ((b[pos + 1] & 0xFF) << 8);
pos += 2 + xlen;
}

if ((flg & 0x08) != 0) {
while (b[pos] != 0) {
pos++;
}
pos++;
}

if ((flg & 0x10) != 0) {
while (b[pos] != 0) {
pos++;
}
pos++;
}

if ((flg & 0x02) != 0) {
pos += 2;
}
return pos;
}

}

static class BodyWriter {
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,191 @@
/*
* Licensed to the Apache Software Foundation (ASF) under one or more
* contributor license agreements. See the NOTICE file distributed with
* this work for additional information regarding copyright ownership.
* The ASF licenses this file to You under the Apache License, Version 2.0
* (the "License"); you may not use this file except in compliance with
* the License. You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/

package org.apache.shenyu.plugin.ai.token.limiter;

import org.slf4j.Logger;
import org.slf4j.LoggerFactory;

import java.io.ByteArrayOutputStream;
import java.util.zip.DataFormatException;
import java.util.zip.Inflater;

/**
* Streaming gzip decoder for handling cross-buffer gzip decompression.
* Package-visible for testing.
*/
class GzipStreamDecoder {

private static final Logger LOG = LoggerFactory.getLogger(GzipStreamDecoder.class);

private final Inflater inflater = new Inflater(true);

private final byte[] decompressBuffer = new byte[4096];

private final GzipHeaderState headerState = new GzipHeaderState();

private boolean abandoned;

/**
* Decode a chunk of gzip data.
* @param inBytes compressed input bytes
* @return decompressed bytes, or empty array if header incomplete or abandoned
*/
byte[] decode(final byte[] inBytes) {
if (abandoned) {
return new byte[0];
}

int offset = 0;
if (!headerState.isComplete()) {
offset = headerState.process(inBytes);
if (headerState.isCapacityExceeded()) {
abandoned = true;
return new byte[0];
}
if (!headerState.isComplete()) {
return new byte[0];
}
}

inflater.setInput(inBytes, offset, inBytes.length - offset);
ByteArrayOutputStream baos = new ByteArrayOutputStream();
try {
int cnt;
while (!inflater.needsInput() && (cnt = inflater.inflate(decompressBuffer)) > 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;
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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);
});
}
Expand Down
Loading
Loading