mirror of
https://gitee.com/mateos/mateclaw.git
synced 2026-09-13 03:13:41 +08:00
feat(llm): 本地模型上下文窗口探测——Ollama/vLLM 真实窗口接入会话预算,超限报错自动反解窗口值
This commit is contained in:
parent
17e220534c
commit
67fd74f7fb
@ -39,6 +39,7 @@ import vip.mate.llm.model.ModelConfigEntity;
|
||||
import vip.mate.llm.model.ModelFamily;
|
||||
import vip.mate.llm.model.ModelProtocol;
|
||||
import vip.mate.llm.model.ModelProviderEntity;
|
||||
import vip.mate.llm.probe.ModelContextWindowResolver;
|
||||
import vip.mate.llm.routing.ProviderModelRef;
|
||||
import vip.mate.llm.routing.ProviderRouter;
|
||||
import vip.mate.llm.service.ModelConfigService;
|
||||
@ -97,6 +98,7 @@ public class AgentGraphBuilder {
|
||||
private final ConversationService conversationService;
|
||||
private final ModelConfigService modelConfigService;
|
||||
private final ModelProviderService modelProviderService;
|
||||
private final ModelContextWindowResolver contextWindowResolver;
|
||||
private final vip.mate.llm.service.ModelCapabilityService modelCapabilityService;
|
||||
private final ProviderRouter providerRouter;
|
||||
private final PlanningService planningService;
|
||||
@ -354,6 +356,12 @@ public class AgentGraphBuilder {
|
||||
|
||||
ModelProtocol protocol = ModelProtocol.fromChatModel(provider.getChatModel());
|
||||
|
||||
// Effective context window: explicit config > local-server probe > null
|
||||
// (downstream keeps its global-default fallback). Without probing, a
|
||||
// local 8k/16k model with maxInputTokens unset budgets against the
|
||||
// 128k global default and the first oversized request fails outright.
|
||||
Integer effectiveMaxInputTokens = contextWindowResolver.resolveMaxInputTokens(provider, runtimeModel);
|
||||
|
||||
// 内置搜索检测(DashScope / Kimi),但不再移除 WebSearchTool — 两者协同而非互斥
|
||||
boolean builtinSearchEnabled = false;
|
||||
Map<String, Object> providerKwargs = modelProviderService.readProviderGenerateKwargs(provider);
|
||||
@ -395,7 +403,7 @@ public class AgentGraphBuilder {
|
||||
// turn by the reasoning / step-execution nodes with the skills loaded
|
||||
// so far this run so load_skill pins float to the top of the catalog.
|
||||
SkillCatalogRenderer skillCatalogRenderer = buildSkillCatalogRenderer(
|
||||
entity, boundTools, runtimeModel.getMaxInputTokens());
|
||||
entity, boundTools, effectiveMaxInputTokens);
|
||||
|
||||
// Extension-tool catalog — only for ReAct. The dynamic tool split runs
|
||||
// in ReasoningNode; Plan-Execute keeps advertising every tool (it has no
|
||||
@ -404,7 +412,7 @@ public class AgentGraphBuilder {
|
||||
boolean isPlanExecute = "plan_execute".equals(entity.getAgentType());
|
||||
if (!isPlanExecute) {
|
||||
String extensionCatalog = toolDisclosureService.renderExtensionCatalog(
|
||||
toolSet, runtimeModel.getMaxInputTokens());
|
||||
toolSet, effectiveMaxInputTokens);
|
||||
if (extensionCatalog != null && !extensionCatalog.isBlank()) {
|
||||
enhancedPrompt = enhancedPrompt + extensionCatalog;
|
||||
}
|
||||
@ -452,7 +460,7 @@ public class AgentGraphBuilder {
|
||||
agent.userLocale = resolveLocale();
|
||||
agent.temperature = runtimeModel.getTemperature();
|
||||
agent.maxTokens = runtimeModel.getMaxTokens();
|
||||
agent.maxInputTokens = runtimeModel.getMaxInputTokens();
|
||||
agent.maxInputTokens = effectiveMaxInputTokens;
|
||||
agent.topP = runtimeModel.getTopP();
|
||||
agent.toolCallingEnabled = toolCallingEnabled;
|
||||
|
||||
@ -566,6 +574,14 @@ public class AgentGraphBuilder {
|
||||
streamTracker, fallbackChain, llmCacheMetricsAggregator, providerHealthTracker,
|
||||
primaryModelConfig != null ? primaryModelConfig.getProvider() : null,
|
||||
providerPool);
|
||||
if (primaryModelConfig != null) {
|
||||
// Feed "prompt too long" rejections back into the window resolver
|
||||
// so the next turn budgets against the server-reported limit.
|
||||
streamingHelper.setContextLimitObserver(errorMessage ->
|
||||
contextWindowResolver.noteContextLimitError(
|
||||
primaryModelConfig.getProvider(),
|
||||
primaryModelConfig.getModelName(), errorMessage));
|
||||
}
|
||||
ToolExecutionExecutor executor = new ToolExecutionExecutor(
|
||||
toolSet, toolGuardService, approvalService, streamTracker,
|
||||
toolTimeoutProperties, toolResultStorage, toolConcurrencyRegistry,
|
||||
@ -839,6 +855,14 @@ public class AgentGraphBuilder {
|
||||
streamTracker, fallbackChain, llmCacheMetricsAggregator, providerHealthTracker,
|
||||
primaryModelConfig != null ? primaryModelConfig.getProvider() : null,
|
||||
providerPool);
|
||||
if (primaryModelConfig != null) {
|
||||
// Feed "prompt too long" rejections back into the window resolver
|
||||
// so the next turn budgets against the server-reported limit.
|
||||
streamingHelper.setContextLimitObserver(errorMessage ->
|
||||
contextWindowResolver.noteContextLimitError(
|
||||
primaryModelConfig.getProvider(),
|
||||
primaryModelConfig.getModelName(), errorMessage));
|
||||
}
|
||||
ToolExecutionExecutor executor = new ToolExecutionExecutor(
|
||||
toolSet, toolGuardService, approvalService, streamTracker,
|
||||
toolTimeoutProperties, toolResultStorage, toolConcurrencyRegistry,
|
||||
|
||||
@ -25,6 +25,7 @@ import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicBoolean;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
import java.util.function.Consumer;
|
||||
|
||||
/**
|
||||
* 节点级流式 LLM 调用辅助
|
||||
@ -162,6 +163,19 @@ public class NodeStreamingChatHelper {
|
||||
this.providerPool = providerPool;
|
||||
}
|
||||
|
||||
/**
|
||||
* Optional hook fired with the raw error chain whenever the PRIMARY model
|
||||
* rejects a call for exceeding its context window. Lets the caller feed
|
||||
* the server-reported limit back into the context-window resolver so the
|
||||
* next turn budgets against the model's true window. Fallback-model
|
||||
* rejections are not reported — they belong to a different model.
|
||||
*/
|
||||
private Consumer<String> contextLimitObserver;
|
||||
|
||||
public void setContextLimitObserver(Consumer<String> observer) {
|
||||
this.contextLimitObserver = observer;
|
||||
}
|
||||
|
||||
private static List<vip.mate.llm.failover.FallbackEntry> wrap(ChatModel m) {
|
||||
// Legacy single-fallback path: providerId is unknown so health tracking
|
||||
// is silently disabled for that one entry (it gets a synthetic id).
|
||||
@ -641,7 +655,7 @@ public class NodeStreamingChatHelper {
|
||||
}
|
||||
llmCallCount++;
|
||||
if (attempt > 0) retryCount++;
|
||||
lastResult = doStreamCall(chatModel, prompt, conversationId, phase, broadcast, attempt);
|
||||
lastResult = doStreamCall(chatModel, prompt, conversationId, phase, broadcast, attempt, true);
|
||||
if (lastResult != null) {
|
||||
// PTL: 不重试,直接返回给上层 Node 处理
|
||||
if (lastResult.errorType() == ErrorType.PROMPT_TOO_LONG) {
|
||||
@ -783,7 +797,7 @@ public class NodeStreamingChatHelper {
|
||||
failoverCount++;
|
||||
llmCallCount++;
|
||||
StreamResult fallbackResult = doStreamCall(fallback, prompt, conversationId,
|
||||
phase + "_fallback_" + (i + 1), broadcast, 0);
|
||||
phase + "_fallback_" + (i + 1), broadcast, 0, false);
|
||||
// Accept only fully successful fallbacks. Non-successful results (auth
|
||||
// error, client error, still-rate-limited) propagate to the next
|
||||
// fallback instead of being surfaced as the final result.
|
||||
@ -830,7 +844,7 @@ public class NodeStreamingChatHelper {
|
||||
*/
|
||||
private StreamResult doStreamCall(ChatModel chatModel, Prompt prompt,
|
||||
String conversationId, String phase,
|
||||
boolean broadcast, int attempt) {
|
||||
boolean broadcast, int attempt, boolean primaryCall) {
|
||||
// Collapse every SystemMessage in the prompt into a single SystemMessage
|
||||
// at index 0. Some OpenAI-compatible providers (LM Studio's built-in
|
||||
// server, certain strict vLLM / SGLang deployments) reject 400
|
||||
@ -883,7 +897,7 @@ public class NodeStreamingChatHelper {
|
||||
}
|
||||
|
||||
try {
|
||||
return doStreamCallInner(chatModel, outbound, conversationId, phase, broadcast, attempt);
|
||||
return doStreamCallInner(chatModel, outbound, conversationId, phase, broadcast, attempt, primaryCall);
|
||||
} finally {
|
||||
// Idempotent: if consumer already took the entry, discard is a no-op.
|
||||
if (relayToken != null) {
|
||||
@ -915,7 +929,7 @@ public class NodeStreamingChatHelper {
|
||||
|
||||
private StreamResult doStreamCallInner(ChatModel chatModel, Prompt prompt,
|
||||
String conversationId, String phase,
|
||||
boolean broadcast, int attempt) {
|
||||
boolean broadcast, int attempt, boolean primaryCall) {
|
||||
if (attempt > 0) {
|
||||
long delay = Math.min(backoffBaseMs * (1L << (attempt - 1)), backoffCapMs);
|
||||
// 加入 jitter 防止雷群效应
|
||||
@ -1246,6 +1260,17 @@ public class NodeStreamingChatHelper {
|
||||
if (errorType == ErrorType.PROMPT_TOO_LONG) {
|
||||
log.warn("[{}] Prompt too long error, returning to node for compaction: {}",
|
||||
phase, error.getMessage());
|
||||
// Teach the context-window resolver the server-reported limit so
|
||||
// the next turn budgets against the model's true window. Raw
|
||||
// chain (incl. response body) — the friendly text may drop the
|
||||
// numbers. Primary model only; fallbacks are different models.
|
||||
if (primaryCall && contextLimitObserver != null) {
|
||||
try {
|
||||
contextLimitObserver.accept(extractFullErrorChain(error));
|
||||
} catch (Exception observerError) {
|
||||
log.debug("context-limit observer failed: {}", observerError.getMessage());
|
||||
}
|
||||
}
|
||||
return buildErrorResultWithType("Prompt 过长: " + extractUserFriendlyError(error),
|
||||
conversationId, phase, errorType);
|
||||
}
|
||||
|
||||
@ -0,0 +1,64 @@
|
||||
package vip.mate.llm.probe;
|
||||
|
||||
import java.util.OptionalInt;
|
||||
import java.util.regex.Matcher;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
/**
|
||||
* Extracts the model's context-window size from a "prompt too long" error
|
||||
* message. Serves as the reconciliation fallback when probing is unavailable:
|
||||
* the serving stack itself states its limit in the rejection text (e.g. vLLM
|
||||
* reports {@code max_model_len}), so one failed call teaches the resolver the
|
||||
* true window for every subsequent turn.
|
||||
*/
|
||||
public final class ContextLimitErrorParser {
|
||||
|
||||
/** Reject absurd parses — anything below one model page or above 10M tokens. */
|
||||
private static final int MIN_PLAUSIBLE = 512;
|
||||
private static final int MAX_PLAUSIBLE = 10_000_000;
|
||||
|
||||
/**
|
||||
* Ordered from most specific to most generic. Each pattern anchors the
|
||||
* number on the limit-keyword side so "requested 50000 tokens, maximum
|
||||
* context length is 32768" yields 32768, not 50000.
|
||||
*/
|
||||
private static final Pattern[] LIMIT_PATTERNS = {
|
||||
// vLLM: "... exceeds the max_model_len 32768" / "max_model_len=32768"
|
||||
Pattern.compile("max_model_len\\D{0,20}?(\\d{3,8})", Pattern.CASE_INSENSITIVE),
|
||||
// OpenAI-style: "This model's maximum context length is 4096 tokens"
|
||||
Pattern.compile("maximum context length is\\s*(\\d{3,8})", Pattern.CASE_INSENSITIVE),
|
||||
// vLLM alt: "maximum model length 32768"
|
||||
Pattern.compile("maximum model length\\D{0,20}?(\\d{3,8})", Pattern.CASE_INSENSITIVE),
|
||||
// Ollama-style knob in the rejection text: "num_ctx 8192"
|
||||
Pattern.compile("num_ctx\\D{0,10}?(\\d{3,8})", Pattern.CASE_INSENSITIVE),
|
||||
// Generic: "context length of only 8192" / "context length limit: 8192"
|
||||
Pattern.compile("context length (?:of only|limit)\\D{0,10}?(\\d{3,8})", Pattern.CASE_INSENSITIVE),
|
||||
};
|
||||
|
||||
private ContextLimitErrorParser() {
|
||||
}
|
||||
|
||||
/**
|
||||
* @return the context window the server reported in the error text, or
|
||||
* empty when no pattern matches or the number is implausible.
|
||||
*/
|
||||
public static OptionalInt extractLimit(String errorMessage) {
|
||||
if (errorMessage == null || errorMessage.isBlank()) {
|
||||
return OptionalInt.empty();
|
||||
}
|
||||
for (Pattern pattern : LIMIT_PATTERNS) {
|
||||
Matcher matcher = pattern.matcher(errorMessage);
|
||||
if (matcher.find()) {
|
||||
try {
|
||||
int value = Integer.parseInt(matcher.group(1));
|
||||
if (value >= MIN_PLAUSIBLE && value <= MAX_PLAUSIBLE) {
|
||||
return OptionalInt.of(value);
|
||||
}
|
||||
} catch (NumberFormatException ignored) {
|
||||
// fall through to the next pattern
|
||||
}
|
||||
}
|
||||
}
|
||||
return OptionalInt.empty();
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,25 @@
|
||||
package vip.mate.llm.probe;
|
||||
|
||||
import lombok.Data;
|
||||
import org.springframework.boot.context.properties.ConfigurationProperties;
|
||||
|
||||
/**
|
||||
* Configuration for local-model context-window probing.
|
||||
*/
|
||||
@Data
|
||||
@ConfigurationProperties(prefix = "mateclaw.context.probe")
|
||||
public class ContextProbeProperties {
|
||||
|
||||
/** Master switch. When false, {@code resolveMaxInputTokens} only honors explicit config. */
|
||||
private boolean enabled = true;
|
||||
|
||||
/** Per-request read timeout. Probing must never hold up chat startup. */
|
||||
private int timeoutMs = 1000;
|
||||
|
||||
/**
|
||||
* How long a probe result (positive or negative) stays cached. Local
|
||||
* servers like LM Studio allow hot-swapping models, so results must not
|
||||
* be persisted — a short in-memory TTL keeps them honest.
|
||||
*/
|
||||
private int cacheTtlSeconds = 600;
|
||||
}
|
||||
@ -0,0 +1,41 @@
|
||||
package vip.mate.llm.probe;
|
||||
|
||||
import vip.mate.llm.model.ModelConfigEntity;
|
||||
import vip.mate.llm.model.ModelProviderEntity;
|
||||
|
||||
import java.util.Optional;
|
||||
|
||||
/**
|
||||
* SPI for probing the real context-window size of a locally hosted model
|
||||
* (Ollama, vLLM, LM Studio, MLX and other self-hosted OpenAI-compatible
|
||||
* servers).
|
||||
*
|
||||
* <p>Motivation: {@code ModelConfigEntity.maxInputTokens} is optional and
|
||||
* rarely filled in for local deployments, so the conversation window budget
|
||||
* silently falls back to the global default (128k). A local 8k/16k model then
|
||||
* never triggers any trimming and the first oversized request fails. Probing
|
||||
* the serving endpoint recovers the true window without user configuration.
|
||||
*
|
||||
* <p>Contract: implementations must be cheap to call (single short HTTP
|
||||
* request), must never throw for routine failures (return
|
||||
* {@link Optional#empty()} instead), and must not be invoked for cloud
|
||||
* providers — {@link #supports} gates that.
|
||||
*/
|
||||
public interface LocalContextProbe {
|
||||
|
||||
/**
|
||||
* @return true when this probe knows how to query the given provider.
|
||||
* Implementations must return false for cloud providers so no
|
||||
* probe traffic ever leaves the local network.
|
||||
*/
|
||||
boolean supports(ModelProviderEntity provider, ModelConfigEntity model);
|
||||
|
||||
/**
|
||||
* Query the serving endpoint for the model's maximum context length.
|
||||
*
|
||||
* @return the context window in tokens, or empty when the endpoint is
|
||||
* unreachable, the model is unknown, or the response carries no
|
||||
* usable length field.
|
||||
*/
|
||||
Optional<Integer> probeContextLength(ModelProviderEntity provider, ModelConfigEntity model);
|
||||
}
|
||||
@ -0,0 +1,85 @@
|
||||
package vip.mate.llm.probe;
|
||||
|
||||
import java.net.InetAddress;
|
||||
import java.net.URI;
|
||||
|
||||
/**
|
||||
* Heuristics for deciding whether a base URL points at a locally hosted /
|
||||
* self-hosted inference server. Probing is restricted to such endpoints so no
|
||||
* probe traffic ever reaches a cloud provider.
|
||||
*/
|
||||
final class LocalEndpoints {
|
||||
|
||||
private LocalEndpoints() {
|
||||
}
|
||||
|
||||
/**
|
||||
* @return true when the URL's host is loopback, a private / link-local
|
||||
* IPv4 range, an mDNS {@code .local} name, or a well-known
|
||||
* container-host alias.
|
||||
*/
|
||||
static boolean isLocal(String baseUrl) {
|
||||
if (baseUrl == null || baseUrl.isBlank()) {
|
||||
return false;
|
||||
}
|
||||
String host;
|
||||
try {
|
||||
host = URI.create(baseUrl.trim()).getHost();
|
||||
} catch (IllegalArgumentException e) {
|
||||
return false;
|
||||
}
|
||||
if (host == null || host.isBlank()) {
|
||||
return false;
|
||||
}
|
||||
String lower = host.toLowerCase();
|
||||
if (lower.equals("localhost") || lower.endsWith(".local")
|
||||
|| lower.equals("host.docker.internal") || lower.equals("host.containers.internal")) {
|
||||
return true;
|
||||
}
|
||||
// Literal IP addresses only — never resolve DNS here: a probe gate
|
||||
// must not add name-resolution latency or leak lookups for cloud hosts.
|
||||
byte[] addr = parseLiteralAddress(lower);
|
||||
if (addr == null) {
|
||||
return false;
|
||||
}
|
||||
try {
|
||||
InetAddress inet = InetAddress.getByAddress(addr);
|
||||
return inet.isLoopbackAddress() || inet.isSiteLocalAddress() || inet.isLinkLocalAddress();
|
||||
} catch (Exception e) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
/** Parse an IPv4/IPv6 literal without triggering DNS. Returns null for hostnames. */
|
||||
private static byte[] parseLiteralAddress(String host) {
|
||||
String h = host;
|
||||
if (h.startsWith("[") && h.endsWith("]")) {
|
||||
h = h.substring(1, h.length() - 1);
|
||||
}
|
||||
if (h.contains(":")) {
|
||||
// IPv6 literal — only loopback matters in practice for local servers.
|
||||
try {
|
||||
return InetAddress.getByName(h).getAddress();
|
||||
} catch (Exception e) {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
String[] parts = h.split("\\.");
|
||||
if (parts.length != 4) {
|
||||
return null;
|
||||
}
|
||||
byte[] out = new byte[4];
|
||||
for (int i = 0; i < 4; i++) {
|
||||
try {
|
||||
int v = Integer.parseInt(parts[i]);
|
||||
if (v < 0 || v > 255) {
|
||||
return null;
|
||||
}
|
||||
out[i] = (byte) v;
|
||||
} catch (NumberFormatException e) {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
return out;
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,125 @@
|
||||
package vip.mate.llm.probe;
|
||||
|
||||
import lombok.RequiredArgsConstructor;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.springframework.boot.context.properties.EnableConfigurationProperties;
|
||||
import org.springframework.stereotype.Service;
|
||||
import vip.mate.llm.model.ModelConfigEntity;
|
||||
import vip.mate.llm.model.ModelProviderEntity;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.OptionalInt;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
|
||||
/**
|
||||
* Resolves the effective context window (max input tokens) for a runtime
|
||||
* model, so downstream window budgeting works from the model's real limit
|
||||
* instead of the 128k global default.
|
||||
*
|
||||
* <p>Priority:
|
||||
* <ol>
|
||||
* <li>explicit {@code ModelConfigEntity.maxInputTokens} — user configuration
|
||||
* always wins;</li>
|
||||
* <li>a probed value from a {@link LocalContextProbe} (runtime-cached with a
|
||||
* short TTL, never persisted — local servers hot-swap models);</li>
|
||||
* <li>{@code null} — caller falls back to the global default, exactly the
|
||||
* pre-probe behavior.</li>
|
||||
* </ol>
|
||||
*
|
||||
* <p>Reconciliation: when a provider rejects a request for being over the
|
||||
* context limit, {@link #noteContextLimitError} parses the limit out of the
|
||||
* error text and seeds the same cache, so the very next turn budgets against
|
||||
* the true window even where probing is unsupported.
|
||||
*/
|
||||
@Slf4j
|
||||
@Service
|
||||
@RequiredArgsConstructor
|
||||
@EnableConfigurationProperties(ContextProbeProperties.class)
|
||||
public class ModelContextWindowResolver {
|
||||
|
||||
private record CacheEntry(Integer value, long expiresAtMs) {
|
||||
}
|
||||
|
||||
private final List<LocalContextProbe> probes;
|
||||
private final ContextProbeProperties properties;
|
||||
|
||||
/** Key: providerId + "/" + modelName. Value may hold null (negative cache). */
|
||||
private final Map<String, CacheEntry> cache = new ConcurrentHashMap<>();
|
||||
|
||||
/**
|
||||
* @return the effective max input tokens, or {@code null} when neither
|
||||
* explicit config nor probing yields a value (caller keeps its
|
||||
* existing global-default fallback).
|
||||
*/
|
||||
public Integer resolveMaxInputTokens(ModelProviderEntity provider, ModelConfigEntity model) {
|
||||
if (model == null) {
|
||||
return null;
|
||||
}
|
||||
if (model.getMaxInputTokens() != null && model.getMaxInputTokens() > 0) {
|
||||
return model.getMaxInputTokens();
|
||||
}
|
||||
if (!properties.isEnabled()) {
|
||||
return null;
|
||||
}
|
||||
String key = cacheKey(provider != null ? provider.getProviderId() : null, model.getModelName());
|
||||
CacheEntry cached = cache.get(key);
|
||||
long now = System.currentTimeMillis();
|
||||
if (cached != null && cached.expiresAtMs() > now) {
|
||||
return cached.value();
|
||||
}
|
||||
Integer probed = null;
|
||||
for (LocalContextProbe probe : probes) {
|
||||
try {
|
||||
if (!probe.supports(provider, model)) {
|
||||
continue;
|
||||
}
|
||||
probed = probe.probeContextLength(provider, model).orElse(null);
|
||||
if (probed != null) {
|
||||
break;
|
||||
}
|
||||
} catch (Exception e) {
|
||||
log.debug("[ContextProbe] probe {} threw for {}: {}",
|
||||
probe.getClass().getSimpleName(), key, e.getMessage());
|
||||
}
|
||||
}
|
||||
cache.put(key, new CacheEntry(probed, now + ttlMs()));
|
||||
if (probed != null) {
|
||||
log.info("[ContextProbe] 探测到模型 {} 的上下文窗口为 {} tokens(未配置 maxInputTokens,窗口预算将使用探测值)",
|
||||
key, probed);
|
||||
}
|
||||
return probed;
|
||||
}
|
||||
|
||||
/**
|
||||
* Feed a "prompt too long" rejection back into the cache. The parsed limit
|
||||
* only takes effect for models without explicit configuration, because
|
||||
* {@link #resolveMaxInputTokens} checks explicit config first.
|
||||
*/
|
||||
public void noteContextLimitError(String providerId, String modelName, String errorMessage) {
|
||||
if (!properties.isEnabled() || modelName == null || modelName.isBlank()) {
|
||||
return;
|
||||
}
|
||||
OptionalInt parsed = ContextLimitErrorParser.extractLimit(errorMessage);
|
||||
if (parsed.isEmpty()) {
|
||||
return;
|
||||
}
|
||||
String key = cacheKey(providerId, modelName);
|
||||
int value = parsed.getAsInt();
|
||||
cache.put(key, new CacheEntry(value, System.currentTimeMillis() + ttlMs()));
|
||||
log.info("[ContextProbe] 从上下文超限报错中解析到模型 {} 的窗口为 {} tokens,已记入运行期缓存", key, value);
|
||||
}
|
||||
|
||||
/** Test hook: drop all cached probe results. */
|
||||
void clearCache() {
|
||||
cache.clear();
|
||||
}
|
||||
|
||||
private long ttlMs() {
|
||||
return Math.max(1, properties.getCacheTtlSeconds()) * 1000L;
|
||||
}
|
||||
|
||||
private static String cacheKey(String providerId, String modelName) {
|
||||
return (providerId == null ? "" : providerId) + "/" + (modelName == null ? "" : modelName);
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,141 @@
|
||||
package vip.mate.llm.probe;
|
||||
|
||||
import com.fasterxml.jackson.databind.JsonNode;
|
||||
import com.fasterxml.jackson.databind.ObjectMapper;
|
||||
import lombok.RequiredArgsConstructor;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.springframework.http.HttpHeaders;
|
||||
import org.springframework.http.MediaType;
|
||||
import org.springframework.http.client.JdkClientHttpRequestFactory;
|
||||
import org.springframework.stereotype.Component;
|
||||
import org.springframework.web.client.RestClient;
|
||||
import vip.mate.llm.model.ModelConfigEntity;
|
||||
import vip.mate.llm.model.ModelProviderEntity;
|
||||
|
||||
import java.net.http.HttpClient;
|
||||
import java.time.Duration;
|
||||
import java.util.Iterator;
|
||||
import java.util.Map;
|
||||
import java.util.Optional;
|
||||
import java.util.OptionalInt;
|
||||
import java.util.regex.Matcher;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
/**
|
||||
* Context-window probe for Ollama servers via the native model-metadata
|
||||
* endpoint ({@code POST /api/show}).
|
||||
*
|
||||
* <p>Resolution order within the response:
|
||||
* <ol>
|
||||
* <li>{@code num_ctx} from the modelfile parameters — the window the server
|
||||
* actually serves with;</li>
|
||||
* <li>the architecture's {@code *.context_length} from {@code model_info} —
|
||||
* an upper bound when no explicit {@code num_ctx} is set.</li>
|
||||
* </ol>
|
||||
* Explicit per-model configuration always wins upstream in the resolver; this
|
||||
* probe only fills the gap when the user configured nothing.
|
||||
*/
|
||||
@Slf4j
|
||||
@Component
|
||||
@RequiredArgsConstructor
|
||||
public class OllamaContextProbe implements LocalContextProbe {
|
||||
|
||||
static final String DEFAULT_BASE_URL = "http://127.0.0.1:11434";
|
||||
|
||||
private static final Pattern NUM_CTX_PATTERN = Pattern.compile("num_ctx\\s+(\\d{3,8})");
|
||||
private static final ObjectMapper OBJECT_MAPPER = new ObjectMapper();
|
||||
|
||||
private final ContextProbeProperties properties;
|
||||
|
||||
@Override
|
||||
public boolean supports(ModelProviderEntity provider, ModelConfigEntity model) {
|
||||
return provider != null && model != null
|
||||
&& "ollama".equalsIgnoreCase(provider.getProviderId());
|
||||
}
|
||||
|
||||
@Override
|
||||
public Optional<Integer> probeContextLength(ModelProviderEntity provider, ModelConfigEntity model) {
|
||||
String baseUrl = normalizeBaseUrl(provider.getBaseUrl());
|
||||
try {
|
||||
RestClient client = RestClient.builder()
|
||||
.requestFactory(requestFactory())
|
||||
.baseUrl(baseUrl)
|
||||
.defaultHeader(HttpHeaders.ACCEPT, MediaType.APPLICATION_JSON_VALUE)
|
||||
.build();
|
||||
// Newer Ollama accepts "model", older releases used "name" — send both.
|
||||
String body = client.post()
|
||||
.uri("/api/show")
|
||||
.contentType(MediaType.APPLICATION_JSON)
|
||||
.body(Map.of("model", model.getModelName(), "name", model.getModelName()))
|
||||
.retrieve()
|
||||
.body(String.class);
|
||||
OptionalInt parsed = parseShowResponse(body);
|
||||
return parsed.isPresent() ? Optional.of(parsed.getAsInt()) : Optional.empty();
|
||||
} catch (Exception e) {
|
||||
log.debug("[ContextProbe] Ollama probe failed for {} at {}: {}",
|
||||
model.getModelName(), baseUrl, e.getMessage());
|
||||
return Optional.empty();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Parse an {@code /api/show} response body. Package-private for tests.
|
||||
*/
|
||||
static OptionalInt parseShowResponse(String body) {
|
||||
if (body == null || body.isBlank()) {
|
||||
return OptionalInt.empty();
|
||||
}
|
||||
try {
|
||||
JsonNode root = OBJECT_MAPPER.readTree(body);
|
||||
// Serving-time knob wins: it is what the server actually allocates.
|
||||
Matcher numCtx = NUM_CTX_PATTERN.matcher(root.path("parameters").asText(""));
|
||||
if (numCtx.find()) {
|
||||
int value = Integer.parseInt(numCtx.group(1));
|
||||
if (value > 0) {
|
||||
return OptionalInt.of(value);
|
||||
}
|
||||
}
|
||||
JsonNode modelInfo = root.path("model_info");
|
||||
if (modelInfo.isObject()) {
|
||||
for (Iterator<String> it = modelInfo.fieldNames(); it.hasNext(); ) {
|
||||
String field = it.next();
|
||||
if (field.endsWith(".context_length")) {
|
||||
int value = modelInfo.path(field).asInt(0);
|
||||
if (value > 0) {
|
||||
return OptionalInt.of(value);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} catch (Exception e) {
|
||||
return OptionalInt.empty();
|
||||
}
|
||||
return OptionalInt.empty();
|
||||
}
|
||||
|
||||
private JdkClientHttpRequestFactory requestFactory() {
|
||||
// HTTP/1.1 pinned: Uvicorn-style local stacks reject the h2c upgrade.
|
||||
HttpClient httpClient = HttpClient.newBuilder()
|
||||
.version(HttpClient.Version.HTTP_1_1)
|
||||
.connectTimeout(Duration.ofMillis(properties.getTimeoutMs()))
|
||||
.build();
|
||||
JdkClientHttpRequestFactory factory = new JdkClientHttpRequestFactory(httpClient);
|
||||
factory.setReadTimeout(Duration.ofMillis(properties.getTimeoutMs()));
|
||||
return factory;
|
||||
}
|
||||
|
||||
/** Ollama providers are often saved with the OpenAI-compatible {@code /v1} suffix — strip it. */
|
||||
static String normalizeBaseUrl(String baseUrl) {
|
||||
if (baseUrl == null || baseUrl.isBlank()) {
|
||||
return DEFAULT_BASE_URL;
|
||||
}
|
||||
String normalized = baseUrl.trim();
|
||||
if (normalized.endsWith("/")) {
|
||||
normalized = normalized.substring(0, normalized.length() - 1);
|
||||
}
|
||||
if (normalized.endsWith("/v1")) {
|
||||
normalized = normalized.substring(0, normalized.length() - 3);
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,130 @@
|
||||
package vip.mate.llm.probe;
|
||||
|
||||
import com.fasterxml.jackson.databind.JsonNode;
|
||||
import com.fasterxml.jackson.databind.ObjectMapper;
|
||||
import lombok.RequiredArgsConstructor;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.springframework.http.HttpHeaders;
|
||||
import org.springframework.http.MediaType;
|
||||
import org.springframework.http.client.JdkClientHttpRequestFactory;
|
||||
import org.springframework.stereotype.Component;
|
||||
import org.springframework.web.client.RestClient;
|
||||
import vip.mate.llm.model.ModelConfigEntity;
|
||||
import vip.mate.llm.model.ModelProtocol;
|
||||
import vip.mate.llm.model.ModelProviderEntity;
|
||||
|
||||
import java.net.http.HttpClient;
|
||||
import java.time.Duration;
|
||||
import java.util.Optional;
|
||||
import java.util.OptionalInt;
|
||||
|
||||
/**
|
||||
* Context-window probe for self-hosted OpenAI-compatible servers (vLLM,
|
||||
* LM Studio, llama.cpp server, MLX, …) via {@code GET /v1/models}.
|
||||
*
|
||||
* <p>vLLM exposes {@code max_model_len} per model entry; other stacks expose
|
||||
* {@code context_length} or {@code max_context_length}. Only endpoints whose
|
||||
* host is local / private are probed — {@link LocalEndpoints#isLocal} gates
|
||||
* that, so no probe traffic is ever sent to a cloud provider.
|
||||
*/
|
||||
@Slf4j
|
||||
@Component
|
||||
@RequiredArgsConstructor
|
||||
public class OpenAiCompatibleContextProbe implements LocalContextProbe {
|
||||
|
||||
private static final ObjectMapper OBJECT_MAPPER = new ObjectMapper();
|
||||
|
||||
private final ContextProbeProperties properties;
|
||||
|
||||
@Override
|
||||
public boolean supports(ModelProviderEntity provider, ModelConfigEntity model) {
|
||||
if (provider == null || model == null) {
|
||||
return false;
|
||||
}
|
||||
// Ollama has a richer native endpoint handled by its dedicated probe.
|
||||
if ("ollama".equalsIgnoreCase(provider.getProviderId())) {
|
||||
return false;
|
||||
}
|
||||
if (ModelProtocol.fromChatModel(provider.getChatModel()) != ModelProtocol.OPENAI_COMPATIBLE) {
|
||||
return false;
|
||||
}
|
||||
return LocalEndpoints.isLocal(provider.getBaseUrl());
|
||||
}
|
||||
|
||||
@Override
|
||||
public Optional<Integer> probeContextLength(ModelProviderEntity provider, ModelConfigEntity model) {
|
||||
String baseUrl = normalizeBaseUrl(provider.getBaseUrl());
|
||||
try {
|
||||
RestClient client = RestClient.builder()
|
||||
.requestFactory(requestFactory())
|
||||
.baseUrl(baseUrl)
|
||||
.defaultHeader(HttpHeaders.ACCEPT, MediaType.APPLICATION_JSON_VALUE)
|
||||
.build();
|
||||
RestClient.RequestHeadersSpec<?> spec = client.get().uri("/v1/models");
|
||||
String apiKey = provider.getApiKey();
|
||||
if (apiKey != null && !apiKey.isBlank()) {
|
||||
spec = spec.header(HttpHeaders.AUTHORIZATION, "Bearer " + apiKey.trim());
|
||||
}
|
||||
String body = spec.retrieve().body(String.class);
|
||||
OptionalInt parsed = parseModelsResponse(body, model.getModelName());
|
||||
return parsed.isPresent() ? Optional.of(parsed.getAsInt()) : Optional.empty();
|
||||
} catch (Exception e) {
|
||||
log.debug("[ContextProbe] OpenAI-compatible probe failed for {} at {}: {}",
|
||||
model.getModelName(), baseUrl, e.getMessage());
|
||||
return Optional.empty();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Find the entry matching {@code modelName} in a {@code /v1/models}
|
||||
* response and read its context-length field. Package-private for tests.
|
||||
*/
|
||||
static OptionalInt parseModelsResponse(String body, String modelName) {
|
||||
if (body == null || body.isBlank() || modelName == null || modelName.isBlank()) {
|
||||
return OptionalInt.empty();
|
||||
}
|
||||
try {
|
||||
JsonNode data = OBJECT_MAPPER.readTree(body).path("data");
|
||||
if (!data.isArray()) {
|
||||
return OptionalInt.empty();
|
||||
}
|
||||
for (JsonNode node : data) {
|
||||
if (!modelName.equals(node.path("id").asText(""))) {
|
||||
continue;
|
||||
}
|
||||
for (String field : new String[]{"max_model_len", "context_length", "max_context_length"}) {
|
||||
int value = node.path(field).asInt(0);
|
||||
if (value > 0) {
|
||||
return OptionalInt.of(value);
|
||||
}
|
||||
}
|
||||
return OptionalInt.empty();
|
||||
}
|
||||
} catch (Exception e) {
|
||||
return OptionalInt.empty();
|
||||
}
|
||||
return OptionalInt.empty();
|
||||
}
|
||||
|
||||
private JdkClientHttpRequestFactory requestFactory() {
|
||||
// HTTP/1.1 pinned: Uvicorn-style local stacks reject the h2c upgrade.
|
||||
HttpClient httpClient = HttpClient.newBuilder()
|
||||
.version(HttpClient.Version.HTTP_1_1)
|
||||
.connectTimeout(Duration.ofMillis(properties.getTimeoutMs()))
|
||||
.build();
|
||||
JdkClientHttpRequestFactory factory = new JdkClientHttpRequestFactory(httpClient);
|
||||
factory.setReadTimeout(Duration.ofMillis(properties.getTimeoutMs()));
|
||||
return factory;
|
||||
}
|
||||
|
||||
private static String normalizeBaseUrl(String baseUrl) {
|
||||
String normalized = baseUrl.trim();
|
||||
if (normalized.endsWith("/")) {
|
||||
normalized = normalized.substring(0, normalized.length() - 1);
|
||||
}
|
||||
if (normalized.endsWith("/v1")) {
|
||||
normalized = normalized.substring(0, normalized.length() - 3);
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,62 @@
|
||||
package vip.mate.llm.probe;
|
||||
|
||||
import org.junit.jupiter.api.DisplayName;
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
import java.util.OptionalInt;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
|
||||
/**
|
||||
* Unit tests for {@link ContextLimitErrorParser} — the reconciliation
|
||||
* fallback that learns the model's context window from rejection text.
|
||||
*/
|
||||
class ContextLimitErrorParserTest {
|
||||
|
||||
@Test
|
||||
@DisplayName("vLLM max_model_len rejection yields the limit, not the requested size")
|
||||
void vllmMaxModelLen() {
|
||||
OptionalInt limit = ContextLimitErrorParser.extractLimit(
|
||||
"This request would exceed the max_model_len 32768 (requested 51234 tokens)");
|
||||
assertEquals(OptionalInt.of(32768), limit);
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("OpenAI-style maximum context length message")
|
||||
void openAiStyle() {
|
||||
OptionalInt limit = ContextLimitErrorParser.extractLimit(
|
||||
"This model's maximum context length is 4096 tokens. However, your messages resulted in 9012 tokens.");
|
||||
assertEquals(OptionalInt.of(4096), limit);
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("vLLM alternate wording: maximum model length")
|
||||
void vllmAlternate() {
|
||||
OptionalInt limit = ContextLimitErrorParser.extractLimit(
|
||||
"Input prompt (40000 tokens) is longer than the maximum model length of 16384");
|
||||
assertEquals(OptionalInt.of(16384), limit);
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("num_ctx wording in rejection text")
|
||||
void numCtx() {
|
||||
OptionalInt limit = ContextLimitErrorParser.extractLimit(
|
||||
"prompt exceeds server window (num_ctx 8192)");
|
||||
assertEquals(OptionalInt.of(8192), limit);
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("no pattern → empty")
|
||||
void unrelatedMessage() {
|
||||
assertTrue(ContextLimitErrorParser.extractLimit("connection refused").isEmpty());
|
||||
assertTrue(ContextLimitErrorParser.extractLimit("").isEmpty());
|
||||
assertTrue(ContextLimitErrorParser.extractLimit(null).isEmpty());
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("implausible numbers are rejected")
|
||||
void implausibleNumbers() {
|
||||
// Below one model page — likely a mis-parse.
|
||||
assertTrue(ContextLimitErrorParser.extractLimit("maximum context length is 100 tokens").isEmpty());
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,135 @@
|
||||
package vip.mate.llm.probe;
|
||||
|
||||
import org.junit.jupiter.api.BeforeEach;
|
||||
import org.junit.jupiter.api.DisplayName;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import vip.mate.llm.model.ModelConfigEntity;
|
||||
import vip.mate.llm.model.ModelProviderEntity;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Optional;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
|
||||
/**
|
||||
* Unit tests for {@link ModelContextWindowResolver} — priority order,
|
||||
* caching, disabled flag, and error-text reconciliation.
|
||||
*/
|
||||
class ModelContextWindowResolverTest {
|
||||
|
||||
private ContextProbeProperties properties;
|
||||
private AtomicInteger probeCalls;
|
||||
|
||||
@BeforeEach
|
||||
void setUp() {
|
||||
properties = new ContextProbeProperties();
|
||||
probeCalls = new AtomicInteger();
|
||||
}
|
||||
|
||||
private LocalContextProbe fixedProbe(Integer value) {
|
||||
return new LocalContextProbe() {
|
||||
@Override
|
||||
public boolean supports(ModelProviderEntity provider, ModelConfigEntity model) {
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Optional<Integer> probeContextLength(ModelProviderEntity provider, ModelConfigEntity model) {
|
||||
probeCalls.incrementAndGet();
|
||||
return Optional.ofNullable(value);
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
private static ModelProviderEntity provider(String id) {
|
||||
ModelProviderEntity provider = new ModelProviderEntity();
|
||||
provider.setProviderId(id);
|
||||
return provider;
|
||||
}
|
||||
|
||||
private static ModelConfigEntity model(String name, Integer maxInputTokens) {
|
||||
ModelConfigEntity model = new ModelConfigEntity();
|
||||
model.setModelName(name);
|
||||
model.setMaxInputTokens(maxInputTokens);
|
||||
return model;
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("explicit maxInputTokens always wins — probe never runs")
|
||||
void explicitConfigWins() {
|
||||
ModelContextWindowResolver resolver =
|
||||
new ModelContextWindowResolver(List.of(fixedProbe(16384)), properties);
|
||||
Integer resolved = resolver.resolveMaxInputTokens(provider("ollama"), model("m", 128000));
|
||||
assertEquals(128000, resolved);
|
||||
assertEquals(0, probeCalls.get());
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("no explicit config → probed value used and cached")
|
||||
void probeFillsGapAndCaches() {
|
||||
ModelContextWindowResolver resolver =
|
||||
new ModelContextWindowResolver(List.of(fixedProbe(16384)), properties);
|
||||
assertEquals(16384, resolver.resolveMaxInputTokens(provider("ollama"), model("m", null)));
|
||||
assertEquals(16384, resolver.resolveMaxInputTokens(provider("ollama"), model("m", 0)));
|
||||
assertEquals(1, probeCalls.get(), "second call must hit the cache");
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("probe miss is negative-cached — the endpoint is not hammered")
|
||||
void negativeCache() {
|
||||
ModelContextWindowResolver resolver =
|
||||
new ModelContextWindowResolver(List.of(fixedProbe(null)), properties);
|
||||
assertNull(resolver.resolveMaxInputTokens(provider("ollama"), model("m", null)));
|
||||
assertNull(resolver.resolveMaxInputTokens(provider("ollama"), model("m", null)));
|
||||
assertEquals(1, probeCalls.get());
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("disabled → null without probing")
|
||||
void disabledSkipsProbing() {
|
||||
properties.setEnabled(false);
|
||||
ModelContextWindowResolver resolver =
|
||||
new ModelContextWindowResolver(List.of(fixedProbe(16384)), properties);
|
||||
assertNull(resolver.resolveMaxInputTokens(provider("ollama"), model("m", null)));
|
||||
assertEquals(0, probeCalls.get());
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("a probe that throws is skipped, not fatal")
|
||||
void throwingProbeIsSkipped() {
|
||||
LocalContextProbe throwing = new LocalContextProbe() {
|
||||
@Override
|
||||
public boolean supports(ModelProviderEntity provider, ModelConfigEntity model) {
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Optional<Integer> probeContextLength(ModelProviderEntity provider, ModelConfigEntity model) {
|
||||
throw new IllegalStateException("boom");
|
||||
}
|
||||
};
|
||||
ModelContextWindowResolver resolver =
|
||||
new ModelContextWindowResolver(List.of(throwing, fixedProbe(8192)), properties);
|
||||
assertEquals(8192, resolver.resolveMaxInputTokens(provider("ollama"), model("m", null)));
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("context-limit error text seeds the cache for later turns")
|
||||
void errorTextReconciliation() {
|
||||
ModelContextWindowResolver resolver =
|
||||
new ModelContextWindowResolver(List.of(), properties);
|
||||
resolver.noteContextLimitError("vllm-local", "m",
|
||||
"Input prompt (40000 tokens) exceeds the max_model_len 32768");
|
||||
assertEquals(32768, resolver.resolveMaxInputTokens(provider("vllm-local"), model("m", null)));
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("unparseable error text changes nothing")
|
||||
void unparseableErrorIgnored() {
|
||||
ModelContextWindowResolver resolver =
|
||||
new ModelContextWindowResolver(List.of(), properties);
|
||||
resolver.noteContextLimitError("p", "m", "connection refused");
|
||||
assertNull(resolver.resolveMaxInputTokens(provider("p"), model("m", null)));
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,52 @@
|
||||
package vip.mate.llm.probe;
|
||||
|
||||
import org.junit.jupiter.api.DisplayName;
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
import java.util.OptionalInt;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
|
||||
/**
|
||||
* Parse-level tests for {@link OllamaContextProbe} against captured
|
||||
* {@code /api/show} response shapes — no HTTP involved.
|
||||
*/
|
||||
class OllamaContextProbeParseTest {
|
||||
|
||||
@Test
|
||||
@DisplayName("num_ctx from modelfile parameters wins over architecture context_length")
|
||||
void numCtxWins() {
|
||||
String body = """
|
||||
{"parameters": "num_ctx 8192\\nstop \\"<|im_end|>\\"",
|
||||
"model_info": {"qwen2.context_length": 32768, "qwen2.embedding_length": 3584}}
|
||||
""";
|
||||
assertEquals(OptionalInt.of(8192), OllamaContextProbe.parseShowResponse(body));
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("architecture context_length used when no num_ctx is set")
|
||||
void contextLengthFallback() {
|
||||
String body = """
|
||||
{"parameters": "stop \\"<|im_end|>\\"",
|
||||
"model_info": {"llama.context_length": 131072, "llama.block_count": 32}}
|
||||
""";
|
||||
assertEquals(OptionalInt.of(131072), OllamaContextProbe.parseShowResponse(body));
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("no usable field → empty")
|
||||
void noUsableField() {
|
||||
assertTrue(OllamaContextProbe.parseShowResponse("{\"model_info\": {}}").isEmpty());
|
||||
assertTrue(OllamaContextProbe.parseShowResponse("not json").isEmpty());
|
||||
assertTrue(OllamaContextProbe.parseShowResponse(null).isEmpty());
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("base URL normalization strips trailing slash and /v1, defaults when blank")
|
||||
void baseUrlNormalization() {
|
||||
assertEquals("http://127.0.0.1:11434", OllamaContextProbe.normalizeBaseUrl(null));
|
||||
assertEquals("http://127.0.0.1:11434", OllamaContextProbe.normalizeBaseUrl(" "));
|
||||
assertEquals("http://192.168.1.5:11434", OllamaContextProbe.normalizeBaseUrl("http://192.168.1.5:11434/v1"));
|
||||
assertEquals("http://192.168.1.5:11434", OllamaContextProbe.normalizeBaseUrl("http://192.168.1.5:11434/"));
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,70 @@
|
||||
package vip.mate.llm.probe;
|
||||
|
||||
import org.junit.jupiter.api.DisplayName;
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
import java.util.OptionalInt;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
|
||||
/**
|
||||
* Parse-level tests for {@link OpenAiCompatibleContextProbe} against
|
||||
* {@code /v1/models} response shapes — no HTTP involved.
|
||||
*/
|
||||
class OpenAiCompatibleContextProbeParseTest {
|
||||
|
||||
@Test
|
||||
@DisplayName("vLLM exposes max_model_len per model entry")
|
||||
void vllmMaxModelLen() {
|
||||
String body = """
|
||||
{"object": "list", "data": [
|
||||
{"id": "Qwen/Qwen2.5-7B-Instruct", "object": "model", "max_model_len": 32768},
|
||||
{"id": "other-model", "object": "model", "max_model_len": 4096}
|
||||
]}
|
||||
""";
|
||||
assertEquals(OptionalInt.of(32768),
|
||||
OpenAiCompatibleContextProbe.parseModelsResponse(body, "Qwen/Qwen2.5-7B-Instruct"));
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("context_length / max_context_length variants are read too")
|
||||
void contextLengthVariants() {
|
||||
String contextLength = "{\"data\": [{\"id\": \"m1\", \"context_length\": 16384}]}";
|
||||
assertEquals(OptionalInt.of(16384),
|
||||
OpenAiCompatibleContextProbe.parseModelsResponse(contextLength, "m1"));
|
||||
|
||||
String maxContextLength = "{\"data\": [{\"id\": \"m2\", \"max_context_length\": 8192}]}";
|
||||
assertEquals(OptionalInt.of(8192),
|
||||
OpenAiCompatibleContextProbe.parseModelsResponse(maxContextLength, "m2"));
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("unknown model id or missing fields → empty")
|
||||
void unknownModelOrMissingField() {
|
||||
String body = "{\"data\": [{\"id\": \"m1\", \"max_model_len\": 32768}]}";
|
||||
assertTrue(OpenAiCompatibleContextProbe.parseModelsResponse(body, "not-there").isEmpty());
|
||||
assertTrue(OpenAiCompatibleContextProbe.parseModelsResponse(
|
||||
"{\"data\": [{\"id\": \"m1\"}]}", "m1").isEmpty());
|
||||
assertTrue(OpenAiCompatibleContextProbe.parseModelsResponse("not json", "m1").isEmpty());
|
||||
assertTrue(OpenAiCompatibleContextProbe.parseModelsResponse(null, "m1").isEmpty());
|
||||
}
|
||||
|
||||
@Test
|
||||
@DisplayName("local endpoint heuristic: loopback and private ranges yes, public hosts no")
|
||||
void localEndpointHeuristic() {
|
||||
assertTrue(LocalEndpoints.isLocal("http://localhost:8000"));
|
||||
assertTrue(LocalEndpoints.isLocal("http://127.0.0.1:8000/v1"));
|
||||
assertTrue(LocalEndpoints.isLocal("http://192.168.1.20:1234"));
|
||||
assertTrue(LocalEndpoints.isLocal("http://10.0.0.3:8000"));
|
||||
assertTrue(LocalEndpoints.isLocal("http://172.16.0.9:8000"));
|
||||
assertTrue(LocalEndpoints.isLocal("http://host.docker.internal:11434"));
|
||||
assertTrue(LocalEndpoints.isLocal("http://mymac.local:1234"));
|
||||
|
||||
assertFalse(LocalEndpoints.isLocal("https://api.openai.com/v1"));
|
||||
assertFalse(LocalEndpoints.isLocal("https://dashscope.aliyuncs.com/compatible-mode/v1"));
|
||||
assertFalse(LocalEndpoints.isLocal("http://172.32.0.1:8000")); // outside 172.16/12
|
||||
assertFalse(LocalEndpoints.isLocal(null));
|
||||
assertFalse(LocalEndpoints.isLocal(""));
|
||||
assertFalse(LocalEndpoints.isLocal("not a url"));
|
||||
}
|
||||
}
|
||||
Loading…
Reference in New Issue
Block a user