diff --git a/README.md b/README.md index 3515d701..a8462144 100644 --- a/README.md +++ b/README.md @@ -83,6 +83,25 @@ Web 工具默认关闭。配置 `lypi.web.enabled=true` 后,运行时会注册 OpenAI 兼容适配支持 Responses、Chat Completions、SSE、WebSocket 和 fallback request style。上层收到的是项目内部的 `AssistantStreamEvent`,不需要直接处理供应商原始事件。模型描述中的 context window、最大输出 token、thinking 支持和图片输入支持会影响请求构建与上下文预算。 +启用 `model-discovery` 的 OpenAI 兼容 Provider 会在应用启动时按配置顺序拉取模型列表;第一个非空结果成为该 Provider 的权威模型集合。远端显式能力字段覆盖 `lypi.ai.model-discovery.defaults`,用户配置的静态同名 `models[]` 再以完整模型描述覆盖远端结果;远端没有返回的静态 model ID 不会进入目录。缺失能力字段默认使用 `context-window=256000`、`max-output-tokens=8192`、`supports-thinking=true` 和 `supports-image-input=true`。所有候选端点都没有返回有效模型时,应用会以不含凭据的端点诊断终止启动。 + +TUI 输入无参数 `/model` 会打开启动期模型快照,候选项统一显示为 `provider/model`;使用上下方向键移动,Enter 切换,Esc 取消。选择结果仍写入会话模型变更条目,恢复会话后继续生效。 + +TUI 的 `/login` 可注册 OpenAI-compatible Provider,交互顺序为: + +```text +/login +1. Channel name +2. Base URL +3. Auth key(掩码显示) +``` + +渠道名必须匹配 `[a-z0-9][a-z0-9_-]{0,63}`,成功后模型以 `/` 出现在 `/model`。同名再次登录表示替换该渠道;URL、凭据、模型发现和持久化任一步失败时,已有运行时渠道保持不变。 + +登录固定使用 OpenAI-compatible Chat Completions over SSE。系统会依次探测 `/models` 和 `/model`,仅在至少发现一个可用模型后才保存并注册 Provider;成功后不会自动切换当前会话模型。日常登录只发现模型目录,不逐模型发送收费能力探针;仓库中的显式真实 E2E 会验证 HIGH thinking、工具续轮和图片 Chat Completions。 + +登录数据仅写入受管文件 `/.ly-pi/login-providers.properties`,不会改写用户维护的 `/.ly-pi/application.yml`,也不缓存发现到的模型列表。应用重启时会重新发现模型,并重新应用当前全局默认值和静态同名完整描述覆盖。 + Anthropic 适配负责 Messages 请求、SSE 事件归一化、tool call/result 映射和 usage 合并。当前版本不启用 Anthropic extended thinking:Anthropic 模型的 `supports-thinking` 应保持 `false`,作为默认模型时还需把 `lypi.runtime.thinking-level` 设为 `off`。 ### 资源与记忆 diff --git a/lypi-ai/src/main/java/cn/lypi/ai/DefaultModelRegistry.java b/lypi-ai/src/main/java/cn/lypi/ai/DefaultModelRegistry.java index 499946de..aad0a567 100644 --- a/lypi-ai/src/main/java/cn/lypi/ai/DefaultModelRegistry.java +++ b/lypi-ai/src/main/java/cn/lypi/ai/DefaultModelRegistry.java @@ -2,28 +2,54 @@ import cn.lypi.contracts.model.ModelDescriptor; import cn.lypi.contracts.model.ModelSelection; +import java.util.ArrayList; import java.util.List; import java.util.Objects; import java.util.Optional; +import java.util.concurrent.atomic.AtomicReference; -public final class DefaultModelRegistry implements ModelRegistry { - private final List descriptors; +public final class DefaultModelRegistry implements RuntimeModelRegistry { + private final AtomicReference> descriptors; public DefaultModelRegistry(List descriptors) { - this.descriptors = List.copyOf(Objects.requireNonNull(descriptors, "descriptors")); + this.descriptors = new AtomicReference<>(immutableDescriptors(descriptors)); } @Override public List list() { - return descriptors; + return descriptors.get(); } @Override public Optional find(ModelSelection selection) { Objects.requireNonNull(selection, "selection"); - return descriptors.stream() + List current = descriptors.get(); + return current.stream() .filter(descriptor -> descriptor.provider().equals(selection.provider())) .filter(descriptor -> descriptor.modelId().equals(selection.modelId())) .findFirst(); } + + @Override + public void replaceProvider(String provider, List replacement) { + String requiredProvider = Objects.requireNonNull(provider, "provider"); + List requiredReplacement = immutableDescriptors(replacement); + for (ModelDescriptor descriptor : requiredReplacement) { + if (!requiredProvider.equals(descriptor.provider())) { + throw new IllegalArgumentException("Replacement descriptor provider must match provider."); + } + } + descriptors.updateAndGet(current -> { + List next = new ArrayList<>(current.size() + requiredReplacement.size()); + current.stream() + .filter(descriptor -> !requiredProvider.equals(descriptor.provider())) + .forEach(next::add); + next.addAll(requiredReplacement); + return List.copyOf(next); + }); + } + + private static List immutableDescriptors(List descriptors) { + return List.copyOf(Objects.requireNonNull(descriptors, "descriptors")); + } } diff --git a/lypi-ai/src/main/java/cn/lypi/ai/ProviderAdapterApiProvider.java b/lypi-ai/src/main/java/cn/lypi/ai/ProviderAdapterApiProvider.java index dcfb1b9b..923f6535 100644 --- a/lypi-ai/src/main/java/cn/lypi/ai/ProviderAdapterApiProvider.java +++ b/lypi-ai/src/main/java/cn/lypi/ai/ProviderAdapterApiProvider.java @@ -10,20 +10,19 @@ import cn.lypi.contracts.runtime.AiProviderRuntimePort; import cn.lypi.contracts.runtime.AiStreamOptions; import cn.lypi.contracts.tool.ToolRegistrySnapshot; +import java.util.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.Objects; -import java.util.function.Function; -import java.util.stream.Collectors; +import java.util.concurrent.atomic.AtomicReference; public final class ProviderAdapterApiProvider implements ApiProvider { private final ApiStyle apiStyle; - private final Map adapters; + private final AtomicReference> adapters; public ProviderAdapterApiProvider(ApiStyle apiStyle, List adapters) { this.apiStyle = Objects.requireNonNull(apiStyle, "apiStyle"); - this.adapters = List.copyOf(Objects.requireNonNull(adapters, "adapters")).stream() - .collect(Collectors.toUnmodifiableMap(ProviderAdapter::provider, Function.identity())); + this.adapters = new AtomicReference<>(indexAdapters(adapters)); } @Override @@ -58,7 +57,7 @@ public AssistantEventStream stream( Objects.requireNonNull(descriptor, "descriptor"); Objects.requireNonNull(options, "options"); Objects.requireNonNull(signal, "signal"); - ProviderAdapter adapter = adapters.get(descriptor.provider()); + ProviderAdapter adapter = adapters.get().get(descriptor.provider()); if (adapter == null) { throw new ModelProviderException( "provider.adapter_unavailable", @@ -69,4 +68,24 @@ public AssistantEventStream stream( } return adapter.stream(context, descriptor, tools, options, signal); } + + public void replaceAdapter(ProviderAdapter adapter) { + ProviderAdapter requiredAdapter = Objects.requireNonNull(adapter, "adapter"); + adapters.updateAndGet(current -> { + Map next = new LinkedHashMap<>(current); + next.put(requiredAdapter.provider(), requiredAdapter); + return Map.copyOf(next); + }); + } + + private static Map indexAdapters(List adapters) { + Map indexed = new LinkedHashMap<>(); + for (ProviderAdapter adapter : List.copyOf(Objects.requireNonNull(adapters, "adapters"))) { + ProviderAdapter previous = indexed.putIfAbsent(adapter.provider(), adapter); + if (previous != null) { + throw new IllegalArgumentException("Duplicate provider adapter: " + adapter.provider()); + } + } + return Map.copyOf(indexed); + } } diff --git a/lypi-ai/src/main/java/cn/lypi/ai/RuntimeModelRegistry.java b/lypi-ai/src/main/java/cn/lypi/ai/RuntimeModelRegistry.java new file mode 100644 index 00000000..f3df40e8 --- /dev/null +++ b/lypi-ai/src/main/java/cn/lypi/ai/RuntimeModelRegistry.java @@ -0,0 +1,11 @@ +package cn.lypi.ai; + +import cn.lypi.contracts.model.ModelDescriptor; +import java.util.List; + +/** + * Boot 装配层用于替换一个 provider 的运行时模型快照。 + */ +public interface RuntimeModelRegistry extends ModelRegistry { + void replaceProvider(String provider, List descriptors); +} diff --git a/lypi-ai/src/main/java/cn/lypi/ai/model/DiscoveredModel.java b/lypi-ai/src/main/java/cn/lypi/ai/model/DiscoveredModel.java new file mode 100644 index 00000000..5bfc859e --- /dev/null +++ b/lypi-ai/src/main/java/cn/lypi/ai/model/DiscoveredModel.java @@ -0,0 +1,32 @@ +package cn.lypi.ai.model; + +import java.util.Optional; +import java.util.OptionalInt; + +public record DiscoveredModel( + String modelId, + OptionalInt contextWindow, + OptionalInt maxOutputTokens, + Optional supportsThinking, + Optional supportsImageInput +) { + public DiscoveredModel { + if (modelId == null || modelId.isBlank()) { + throw new IllegalArgumentException("modelId is required"); + } + contextWindow = contextWindow == null ? OptionalInt.empty() : contextWindow; + maxOutputTokens = maxOutputTokens == null ? OptionalInt.empty() : maxOutputTokens; + supportsThinking = supportsThinking == null ? Optional.empty() : supportsThinking; + supportsImageInput = supportsImageInput == null ? Optional.empty() : supportsImageInput; + } + + public static DiscoveredModel idOnly(String modelId) { + return new DiscoveredModel( + modelId, + OptionalInt.empty(), + OptionalInt.empty(), + Optional.empty(), + Optional.empty() + ); + } +} diff --git a/lypi-ai/src/main/java/cn/lypi/ai/model/DiscoveredModelDefaults.java b/lypi-ai/src/main/java/cn/lypi/ai/model/DiscoveredModelDefaults.java new file mode 100644 index 00000000..a2e1f0ca --- /dev/null +++ b/lypi-ai/src/main/java/cn/lypi/ai/model/DiscoveredModelDefaults.java @@ -0,0 +1,27 @@ +package cn.lypi.ai.model; + +import cn.lypi.contracts.model.CostProfile; +import java.util.Map; +import java.util.Objects; + +public record DiscoveredModelDefaults( + int contextWindow, + int maxOutputTokens, + boolean supportsThinking, + boolean supportsImageInput, + CostProfile costProfile, + Map compat +) { + public static final int DEFAULT_CONTEXT_WINDOW = 256_000; + public static final int DEFAULT_MAX_OUTPUT_TOKENS = 8_192; + public static final boolean DEFAULT_SUPPORTS_THINKING = true; + public static final boolean DEFAULT_SUPPORTS_IMAGE_INPUT = true; + + public DiscoveredModelDefaults { + if (contextWindow <= 0 || maxOutputTokens <= 0) { + throw new IllegalArgumentException("Model discovery default token limits must be positive."); + } + costProfile = Objects.requireNonNull(costProfile, "costProfile"); + compat = Map.copyOf(Objects.requireNonNull(compat, "compat")); + } +} diff --git a/lypi-ai/src/main/java/cn/lypi/ai/model/DiscoveredModelDescriptorMapper.java b/lypi-ai/src/main/java/cn/lypi/ai/model/DiscoveredModelDescriptorMapper.java new file mode 100644 index 00000000..93937971 --- /dev/null +++ b/lypi-ai/src/main/java/cn/lypi/ai/model/DiscoveredModelDescriptorMapper.java @@ -0,0 +1,41 @@ +package cn.lypi.ai.model; + +import cn.lypi.contracts.model.ApiStyle; +import cn.lypi.contracts.model.ModelDescriptor; +import java.net.URI; +import java.util.Objects; + +public final class DiscoveredModelDescriptorMapper { + private final String provider; + private final URI baseUrl; + private final ApiStyle apiStyle; + private final DiscoveredModelDefaults defaults; + + public DiscoveredModelDescriptorMapper( + String provider, + URI baseUrl, + ApiStyle apiStyle, + DiscoveredModelDefaults defaults + ) { + this.provider = Objects.requireNonNull(provider, "provider"); + this.baseUrl = Objects.requireNonNull(baseUrl, "baseUrl"); + this.apiStyle = Objects.requireNonNull(apiStyle, "apiStyle"); + this.defaults = Objects.requireNonNull(defaults, "defaults"); + } + + public ModelDescriptor map(DiscoveredModel discovered) { + Objects.requireNonNull(discovered, "discovered"); + return new ModelDescriptor( + provider, + discovered.modelId(), + baseUrl, + apiStyle, + discovered.contextWindow().orElse(defaults.contextWindow()), + discovered.maxOutputTokens().orElse(defaults.maxOutputTokens()), + discovered.supportsThinking().orElse(defaults.supportsThinking()), + discovered.supportsImageInput().orElse(defaults.supportsImageInput()), + defaults.costProfile(), + CompatSanitizer.sanitize(defaults.compat()) + ); + } +} diff --git a/lypi-ai/src/main/java/cn/lypi/ai/model/RemoteModelDescriptorSource.java b/lypi-ai/src/main/java/cn/lypi/ai/model/RemoteModelDescriptorSource.java index 9a29f20a..863a181d 100644 --- a/lypi-ai/src/main/java/cn/lypi/ai/model/RemoteModelDescriptorSource.java +++ b/lypi-ai/src/main/java/cn/lypi/ai/model/RemoteModelDescriptorSource.java @@ -1,25 +1,20 @@ package cn.lypi.ai.model; import cn.lypi.contracts.model.ApiStyle; -import cn.lypi.contracts.model.CostProfile; import cn.lypi.contracts.model.ModelDescriptor; import java.net.URI; import java.time.Duration; -import java.util.ArrayList; import java.util.List; -import java.util.Map; import java.util.Objects; public final class RemoteModelDescriptorSource implements ModelDescriptorSource { private final boolean enabled; - private final String provider; private final URI baseUrl; - private final ApiStyle apiStyle; private final String apiKey; private final List discoveryPaths; private final Duration timeout; private final RemoteModelDiscoveryClient client; - private final DescriptorDefaults defaults; + private final DiscoveredModelDescriptorMapper mapper; public RemoteModelDescriptorSource( boolean enabled, @@ -30,17 +25,15 @@ public RemoteModelDescriptorSource( List discoveryPaths, Duration timeout, RemoteModelDiscoveryClient client, - DescriptorDefaults defaults + DiscoveredModelDefaults defaults ) { this.enabled = enabled; - this.provider = Objects.requireNonNull(provider, "provider"); this.baseUrl = Objects.requireNonNull(baseUrl, "baseUrl"); - this.apiStyle = Objects.requireNonNull(apiStyle, "apiStyle"); this.apiKey = apiKey; this.discoveryPaths = List.copyOf(Objects.requireNonNull(discoveryPaths, "discoveryPaths")); this.timeout = timeout == null ? Duration.ofSeconds(30) : timeout; this.client = Objects.requireNonNull(client, "client"); - this.defaults = Objects.requireNonNull(defaults, "defaults"); + this.mapper = new DiscoveredModelDescriptorMapper(provider, this.baseUrl, apiStyle, defaults); } @Override @@ -48,35 +41,8 @@ public List list() { if (!enabled) { return List.of(); } - List descriptors = new ArrayList<>(); - for (String modelId : client.discover(baseUrl, apiKey, discoveryPaths, timeout)) { - descriptors.add(new ModelDescriptor( - provider, - modelId, - baseUrl, - apiStyle, - defaults.contextWindow(), - defaults.maxOutputTokens(), - defaults.supportsThinking(), - defaults.supportsImageInput(), - defaults.costProfile(), - CompatSanitizer.sanitize(defaults.compat()) - )); - } - return descriptors; - } - - public record DescriptorDefaults( - int contextWindow, - int maxOutputTokens, - boolean supportsThinking, - boolean supportsImageInput, - CostProfile costProfile, - Map compat - ) { - public DescriptorDefaults { - Objects.requireNonNull(costProfile, "costProfile"); - compat = Map.copyOf(Objects.requireNonNull(compat, "compat")); - } + return client.discoverModels(baseUrl, apiKey, discoveryPaths, timeout).stream() + .map(mapper::map) + .toList(); } } diff --git a/lypi-ai/src/main/java/cn/lypi/ai/model/RemoteModelDiscoveryClient.java b/lypi-ai/src/main/java/cn/lypi/ai/model/RemoteModelDiscoveryClient.java index dce41f60..d59e4f9b 100644 --- a/lypi-ai/src/main/java/cn/lypi/ai/model/RemoteModelDiscoveryClient.java +++ b/lypi-ai/src/main/java/cn/lypi/ai/model/RemoteModelDiscoveryClient.java @@ -1,5 +1,7 @@ package cn.lypi.ai.model; +import cn.lypi.contracts.error.ErrorSeverity; +import cn.lypi.contracts.error.ModelProviderException; import com.fasterxml.jackson.databind.JsonNode; import com.fasterxml.jackson.databind.ObjectMapper; import java.io.IOException; @@ -9,8 +11,12 @@ import java.net.http.HttpResponse; import java.time.Duration; import java.util.ArrayList; +import java.util.LinkedHashMap; import java.util.List; +import java.util.Map; import java.util.Objects; +import java.util.Optional; +import java.util.OptionalInt; public class RemoteModelDiscoveryClient { private final HttpClient httpClient; @@ -26,40 +32,84 @@ public RemoteModelDiscoveryClient(HttpClient httpClient, ObjectMapper objectMapp } public List discover(URI baseUrl, String apiKey, List paths, Duration timeout) { + return discoverModels(baseUrl, apiKey, paths, timeout).stream() + .map(DiscoveredModel::modelId) + .toList(); + } + + public List discoverModels( + URI baseUrl, + String apiKey, + List paths, + Duration timeout + ) { Objects.requireNonNull(baseUrl, "baseUrl"); Objects.requireNonNull(paths, "paths"); Duration requestTimeout = timeout == null ? Duration.ofSeconds(30) : timeout; + List diagnostics = new ArrayList<>(); for (String path : paths) { - List modelIds = request(baseUrl, apiKey, path, requestTimeout); - if (!modelIds.isEmpty()) { - return modelIds; + URI endpoint = endpoint(baseUrl, path); + DiscoveryAttempt attempt = request(endpoint, apiKey, requestTimeout); + if (!attempt.models().isEmpty()) { + return attempt.models(); + } + diagnostics.add(safeEndpoint(endpoint) + ": " + attempt.diagnostic()); + if (attempt.interrupted()) { + break; } } - return List.of(); + String details = diagnostics.isEmpty() + ? "no candidate endpoints configured" + : String.join("; ", diagnostics); + throw new ModelProviderException( + "model.discovery_unavailable", + ErrorSeverity.ERROR, + false, + "Remote model discovery returned no usable models. " + details + ); } - private List request(URI baseUrl, String apiKey, String path, Duration timeout) { - HttpRequest.Builder builder = HttpRequest.newBuilder(endpoint(baseUrl, path)) - .timeout(timeout) - .GET(); - if (apiKey != null && !apiKey.isBlank()) { - builder.header("Authorization", "Bearer " + apiKey); - } + private DiscoveryAttempt request(URI endpoint, String apiKey, Duration timeout) { + HttpRequest request; try { - HttpResponse response = httpClient.send(builder.build(), HttpResponse.BodyHandlers.ofString()); - if (response.statusCode() < 200 || response.statusCode() >= 300) { - return List.of(); + HttpRequest.Builder builder = HttpRequest.newBuilder(endpoint) + .timeout(timeout) + .GET(); + if (apiKey != null && !apiKey.isBlank()) { + builder.header("Authorization", "Bearer " + apiKey); } - return parse(response.body()); - } catch (IOException | InterruptedException | RuntimeException error) { - if (error instanceof InterruptedException) { - Thread.currentThread().interrupt(); + request = builder.build(); + } catch (RuntimeException error) { + return DiscoveryAttempt.failure("invalid request configuration"); + } + + HttpResponse response; + try { + response = httpClient.send(request, HttpResponse.BodyHandlers.ofString()); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + return DiscoveryAttempt.interruptedFailure(); + } catch (IOException error) { + return DiscoveryAttempt.failure("network error: " + error.getClass().getSimpleName()); + } catch (RuntimeException error) { + return DiscoveryAttempt.failure("request error: " + error.getClass().getSimpleName()); + } + + if (response.statusCode() < 200 || response.statusCode() >= 300) { + return DiscoveryAttempt.failure("HTTP " + response.statusCode()); + } + try { + List models = parse(response.body()); + if (models.isEmpty()) { + return DiscoveryAttempt.failure("response contained no usable model ids"); } - return List.of(); + return DiscoveryAttempt.success(models); + } catch (IOException | RuntimeException error) { + return DiscoveryAttempt.failure("invalid JSON response"); } } - private List parse(String body) throws IOException { + private List parse(String body) throws IOException { if (body == null || body.isBlank()) { return List.of(); } @@ -67,35 +117,130 @@ private List parse(String body) throws IOException { if (root.isArray()) { return stringArray(root); } - List dataModels = objectArrayIds(root.path("data")); + List dataModels = objectArrayModels(root.path("data")); if (!dataModels.isEmpty()) { return dataModels; } - return objectArrayIds(root.path("models")); + return objectArrayModels(root.path("models")); } - private static List stringArray(JsonNode node) { - List modelIds = new ArrayList<>(); + private static List stringArray(JsonNode node) { + Map models = new LinkedHashMap<>(); for (JsonNode item : node) { if (item.isTextual() && !item.asText().isBlank()) { - modelIds.add(item.asText()); + models.putIfAbsent(item.asText(), DiscoveredModel.idOnly(item.asText())); } } - return List.copyOf(modelIds); + return List.copyOf(models.values()); } - private static List objectArrayIds(JsonNode node) { + private static List objectArrayModels(JsonNode node) { if (!node.isArray()) { return List.of(); } - List modelIds = new ArrayList<>(); + Map models = new LinkedHashMap<>(); for (JsonNode item : node) { JsonNode id = item.path("id"); if (id.isTextual() && !id.asText().isBlank()) { - modelIds.add(id.asText()); + models.putIfAbsent(id.asText(), discoveredModel(id.asText(), item)); + } + } + return List.copyOf(models.values()); + } + + private static DiscoveredModel discoveredModel(String modelId, JsonNode item) { + Optional supportsThinking = firstBoolean( + item.path("supportsThinking"), + item.path("supports_thinking"), + item.path("supportsReasoning"), + item.path("supports_reasoning") + ); + if (supportsThinking.isEmpty() && supportsReasoningParameter(item.path("supported_parameters"))) { + supportsThinking = Optional.of(true); + } + + Optional supportsImageInput = firstBoolean( + item.path("supportsImageInput"), + item.path("supports_image_input") + ); + if (supportsImageInput.isEmpty()) { + supportsImageInput = firstImageCapability( + item.path("input_modalities"), + item.path("architecture").path("input_modalities") + ); + } + + return new DiscoveredModel( + modelId, + firstPositiveInt( + item.path("contextWindow"), + item.path("context_length"), + item.path("context_window") + ), + firstPositiveInt( + item.path("maxOutputTokens"), + item.path("max_output_tokens"), + item.path("max_tokens"), + item.path("top_provider").path("max_completion_tokens") + ), + supportsThinking, + supportsImageInput + ); + } + + private static OptionalInt firstPositiveInt(JsonNode... candidates) { + for (JsonNode candidate : candidates) { + if (candidate.isIntegralNumber() && candidate.canConvertToInt() && candidate.intValue() > 0) { + return OptionalInt.of(candidate.intValue()); + } + } + return OptionalInt.empty(); + } + + private static Optional firstBoolean(JsonNode... candidates) { + for (JsonNode candidate : candidates) { + if (candidate.isBoolean()) { + return Optional.of(candidate.booleanValue()); + } + } + return Optional.empty(); + } + + private static boolean supportsReasoningParameter(JsonNode parameters) { + if (!parameters.isArray()) { + return false; + } + for (JsonNode parameter : parameters) { + if (!parameter.isTextual()) { + continue; + } + String value = parameter.asText(); + if ("reasoning".equals(value) || "reasoning_effort".equals(value) || "thinking".equals(value)) { + return true; + } + } + return false; + } + + private static Optional firstImageCapability(JsonNode... modalitiesCandidates) { + for (JsonNode modalities : modalitiesCandidates) { + if (!modalities.isArray() || modalities.isEmpty()) { + continue; + } + boolean onlyText = true; + for (JsonNode modality : modalities) { + if (modality.isTextual() && "image".equals(modality.asText())) { + return Optional.of(true); + } + if (!modality.isTextual() || !"text".equals(modality.asText())) { + onlyText = false; + } + } + if (onlyText) { + return Optional.of(false); } } - return List.copyOf(modelIds); + return Optional.empty(); } private static URI endpoint(URI baseUrl, String path) { @@ -105,4 +250,32 @@ private static URI endpoint(URI baseUrl, String path) { normalizedPath = normalizedPath.startsWith("/") ? normalizedPath.substring(1) : normalizedPath; return URI.create(normalizedBase + "/" + normalizedPath); } + + private static String safeEndpoint(URI endpoint) { + String path = endpoint.getRawPath() == null || endpoint.getRawPath().isBlank() ? "/" : endpoint.getRawPath(); + if (endpoint.getHost() == null) { + return path; + } + String port = endpoint.getPort() < 0 ? "" : ":" + endpoint.getPort(); + return endpoint.getScheme() + "://" + endpoint.getHost() + port + path; + } + + private record DiscoveryAttempt(List models, String diagnostic, boolean interrupted) { + private DiscoveryAttempt { + models = List.copyOf(models); + diagnostic = diagnostic == null ? "unknown failure" : diagnostic; + } + + private static DiscoveryAttempt success(List models) { + return new DiscoveryAttempt(models, "", false); + } + + private static DiscoveryAttempt failure(String diagnostic) { + return new DiscoveryAttempt(List.of(), diagnostic, false); + } + + private static DiscoveryAttempt interruptedFailure() { + return new DiscoveryAttempt(List.of(), "request interrupted", true); + } + } } diff --git a/lypi-ai/src/main/java/cn/lypi/ai/provider/ProviderRequest.java b/lypi-ai/src/main/java/cn/lypi/ai/provider/ProviderRequest.java index 36df17dc..426e9fce 100644 --- a/lypi-ai/src/main/java/cn/lypi/ai/provider/ProviderRequest.java +++ b/lypi-ai/src/main/java/cn/lypi/ai/provider/ProviderRequest.java @@ -19,4 +19,9 @@ public ProviderRequest(URI uri, Map headers, String body) { headers = Map.copyOf(headers); timeout = timeout == null ? Optional.empty() : timeout; } + + @Override + public String toString() { + return "ProviderRequest[uri=" + uri + ", headers=, body=, timeout=" + timeout + "]"; + } } diff --git a/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsRequestBuilder.java b/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsRequestBuilder.java index d4082b21..86009586 100644 --- a/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsRequestBuilder.java +++ b/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsRequestBuilder.java @@ -22,6 +22,9 @@ import java.util.Optional; public final class OpenAiChatCompletionsRequestBuilder { + private static final String REASONING_CONTENT_COMPAT = + "requires-reasoning-content-on-assistant-messages"; + private final ObjectMapper objectMapper; public OpenAiChatCompletionsRequestBuilder() { @@ -46,7 +49,7 @@ public ObjectNode build(LypiModelRequest request, OpenAiProviderConfig config) { request.options().maxOutputTokens().ifPresent(tokens -> body.put("max_tokens", tokens)); request.options().temperature().ifPresent(temperature -> body.put("temperature", temperature)); promptCacheKey(request).ifPresent(key -> body.put("prompt_cache_key", key)); - body.set("messages", messages(request)); + body.set("messages", messages(request, config)); if (!request.tools().isEmpty()) { body.set("tools", tools(request)); } @@ -54,7 +57,7 @@ public ObjectNode build(LypiModelRequest request, OpenAiProviderConfig config) { return body; } - private ArrayNode messages(LypiModelRequest request) { + private ArrayNode messages(LypiModelRequest request, OpenAiProviderConfig config) { ArrayNode messages = objectMapper.createArrayNode(); if (!request.systemPrompt().isBlank()) { ObjectNode system = objectMapper.createObjectNode(); @@ -62,41 +65,25 @@ private ArrayNode messages(LypiModelRequest request) { system.put("content", request.systemPrompt()); messages.add(system); } + boolean requiresReasoningContent = compatEnabled(config, REASONING_CONTENT_COMPAT); for (LypiMessage message : request.messages()) { - appendMessage(messages, message); + appendMessage(messages, message, requiresReasoningContent); } return messages; } - private void appendMessage(ArrayNode messages, LypiMessage message) { - List toolCallBlocks = message.content().stream() - .filter(LypiToolCallBlock.class::isInstance) - .map(LypiToolCallBlock.class::cast) - .toList(); - if (message.role() == LypiRole.ASSISTANT && !toolCallBlocks.isEmpty()) { - ObjectNode node = objectMapper.createObjectNode(); - node.put("role", "assistant"); - String content = assistantToolCallContent(message); - if (content.isBlank()) { - node.putNull("content"); - } else { - node.put("content", content); - } - ArrayNode toolCalls = objectMapper.createArrayNode(); - for (LypiToolCallBlock block : toolCallBlocks) { - toolCalls.add(toolCall(block)); - } - node.set("tool_calls", toolCalls); - messages.add(node); + private void appendMessage(ArrayNode messages, LypiMessage message, boolean requiresReasoningContent) { + if (message.role() == LypiRole.ASSISTANT) { + messages.add(assistantMessage(message, requiresReasoningContent)); + return; + } + if (hasImageAttachment(message)) { + appendMultimodalMessage(messages, message); return; } for (LypiContentBlock block : message.content()) { if (block instanceof LypiToolResultBlock toolResult) { - ObjectNode node = objectMapper.createObjectNode(); - node.put("role", "tool"); - node.put("tool_call_id", toolResult.toolUseId()); - node.put("content", toolResult.text()); - messages.add(node); + messages.add(toolResultMessage(toolResult)); } else { ObjectNode node = objectMapper.createObjectNode(); node.put("role", role(message.role())); @@ -106,14 +93,125 @@ private void appendMessage(ArrayNode messages, LypiMessage message) { } } - private String assistantToolCallContent(LypiMessage message) { + private void appendMultimodalMessage(ArrayNode messages, LypiMessage message) { + if (message.role() == LypiRole.TOOL_RESULT) { + message.content().stream() + .filter(LypiToolResultBlock.class::isInstance) + .map(LypiToolResultBlock.class::cast) + .map(this::toolResultMessage) + .forEach(messages::add); + } + List contentBlocks = message.content().stream() + .filter(block -> !(block instanceof LypiToolResultBlock)) + .toList(); + if (contentBlocks.isEmpty()) { + return; + } + ObjectNode node = objectMapper.createObjectNode(); + node.put("role", role(message.role())); + node.set("content", multimodalContent(contentBlocks)); + messages.add(node); + } + + private ObjectNode toolResultMessage(LypiToolResultBlock toolResult) { + ObjectNode node = objectMapper.createObjectNode(); + node.put("role", "tool"); + node.put("tool_call_id", toolResult.toolUseId()); + node.put("content", toolResult.text()); + return node; + } + + private ArrayNode multimodalContent(List blocks) { + ArrayNode content = objectMapper.createArrayNode(); + for (LypiContentBlock block : blocks) { + if (block instanceof LypiAttachmentBlock attachment) { + Optional url = imageUrl(attachment); + if (url.isPresent()) { + content.add(imageContent(attachment, url.orElseThrow())); + continue; + } + } + String text = blockText(block); + if (text != null && !text.isBlank()) { + ObjectNode part = objectMapper.createObjectNode(); + part.put("type", "text"); + part.put("text", text); + content.add(part); + } + } + return content; + } + + private ObjectNode imageContent(LypiAttachmentBlock attachment, String url) { + ObjectNode part = objectMapper.createObjectNode(); + part.put("type", "image_url"); + ObjectNode imageUrl = objectMapper.createObjectNode(); + imageUrl.put("url", url); + imageUrl.put("detail", String.valueOf(attachment.metadata().getOrDefault("detail", "high"))); + part.set("image_url", imageUrl); + return part; + } + + private boolean hasImageAttachment(LypiMessage message) { + return message.content().stream() + .filter(LypiAttachmentBlock.class::isInstance) + .map(LypiAttachmentBlock.class::cast) + .anyMatch(attachment -> imageUrl(attachment).isPresent()); + } + + private Optional imageUrl(LypiAttachmentBlock attachment) { + Object rawUrl = attachment.metadata().get("imageUrl"); + if (rawUrl == null) { + return Optional.empty(); + } + String url = String.valueOf(rawUrl); + return url.isBlank() ? Optional.empty() : Optional.of(url); + } + + private ObjectNode assistantMessage(LypiMessage message, boolean requiresReasoningContent) { + List toolCallBlocks = message.content().stream() + .filter(LypiToolCallBlock.class::isInstance) + .map(LypiToolCallBlock.class::cast) + .toList(); + ObjectNode node = objectMapper.createObjectNode(); + node.put("role", "assistant"); + String content = assistantContent(message, requiresReasoningContent); + if (content.isBlank()) { + node.putNull("content"); + } else { + node.put("content", content); + } + if (requiresReasoningContent) { + node.put("reasoning_content", assistantThinking(message)); + } + if (!toolCallBlocks.isEmpty()) { + ArrayNode toolCalls = objectMapper.createArrayNode(); + for (LypiToolCallBlock block : toolCallBlocks) { + toolCalls.add(toolCall(block)); + } + node.set("tool_calls", toolCalls); + } + return node; + } + + private String assistantContent(LypiMessage message, boolean excludesThinking) { return message.content().stream() .filter(block -> !(block instanceof LypiToolCallBlock)) + .filter(block -> !excludesThinking || !(block instanceof LypiThinkingBlock)) .map(this::blockText) .filter(text -> text != null && !text.isBlank()) .collect(java.util.stream.Collectors.joining("\n")); } + private String assistantThinking(LypiMessage message) { + return message.content().stream() + .filter(LypiThinkingBlock.class::isInstance) + .map(LypiThinkingBlock.class::cast) + .map(LypiThinkingBlock::text) + .filter(text -> text != null && !text.isBlank()) + .collect(java.util.stream.Collectors.joining("\n")); + } + private ObjectNode toolCall(LypiToolCallBlock toolCall) { ObjectNode wrapper = objectMapper.createObjectNode(); ObjectNode function = objectMapper.createObjectNode(); @@ -182,4 +280,12 @@ private Optional promptCacheKey(LypiModelRequest request) { String value = String.valueOf(key); return value.isBlank() ? Optional.empty() : Optional.of(value); } + + private boolean compatEnabled(OpenAiProviderConfig config, String key) { + Object value = config.compat().get(key); + if (value instanceof Boolean enabled) { + return enabled; + } + return value instanceof String text && Boolean.parseBoolean(text); + } } diff --git a/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsStreamNormalizer.java b/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsStreamNormalizer.java index b7aea48b..e739a92e 100644 --- a/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsStreamNormalizer.java +++ b/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsStreamNormalizer.java @@ -4,7 +4,6 @@ import cn.lypi.contracts.model.AssistantError; import cn.lypi.contracts.model.AssistantStart; import cn.lypi.contracts.model.AssistantStreamEvent; -import cn.lypi.contracts.model.TextDelta; import cn.lypi.contracts.model.ThinkingDelta; import cn.lypi.contracts.model.TokenUsage; import cn.lypi.contracts.model.ToolCallDelta; @@ -14,11 +13,13 @@ import java.util.ArrayList; import java.util.LinkedHashMap; import java.util.List; +import java.util.Locale; import java.util.Map; import java.util.Optional; public final class OpenAiChatCompletionsStreamNormalizer implements OpenAiStreamNormalizer { private final ObjectMapper objectMapper; + private final TaggedThinkingContentParser taggedThinkingContentParser = new TaggedThinkingContentParser(); private final Map toolCalls = new LinkedHashMap<>(); private boolean started; private boolean doneEmitted; @@ -73,11 +74,17 @@ public List normalize(String data) { return normalized; } JsonNode delta = choice.path("delta"); - if (delta.hasNonNull("content")) { - normalized.add(new TextDelta(delta.path("content").asText())); + Optional reasoning = firstNonBlankText( + delta.get("reasoning_content"), + delta.get("reasoning"), + delta.get("reasoning_text") + ); + if (reasoning.isEmpty()) { + reasoning = firstReasoningDetailText(delta.path("reasoning_details")); } - if (delta.hasNonNull("reasoning_content")) { - normalized.add(new ThinkingDelta(delta.path("reasoning_content").asText())); + reasoning.ifPresent(value -> normalized.add(new ThinkingDelta(value))); + if (delta.path("content").isTextual()) { + normalized.addAll(taggedThinkingContentParser.accept(delta.path("content").asText())); } if (delta.path("tool_calls").isArray()) { delta.path("tool_calls").forEach(toolCall -> normalized.add(toolCallDelta(toolCall))); @@ -121,7 +128,38 @@ private List done(Optional usage) { return List.of(); } doneEmitted = true; - return List.of(new AssistantDone(usage, Optional.of("stop"))); + List events = new ArrayList<>(taggedThinkingContentParser.finish()); + events.add(new AssistantDone(usage, Optional.of("stop"))); + return List.copyOf(events); + } + + private Optional firstNonBlankText(JsonNode... candidates) { + for (JsonNode candidate : candidates) { + if (candidate != null && candidate.isTextual() && !candidate.asText().isBlank()) { + return Optional.of(candidate.asText()); + } + } + return Optional.empty(); + } + + private Optional firstReasoningDetailText(JsonNode details) { + if (!details.isArray()) { + return Optional.empty(); + } + for (JsonNode detail : details) { + if (!detail.isObject()) { + continue; + } + String type = detail.path("type").asText("").toLowerCase(Locale.ROOT); + if (type.contains("encrypted") || type.contains("signature")) { + continue; + } + Optional text = firstNonBlankText(detail.get("text")); + if (text.isPresent()) { + return text; + } + } + return Optional.empty(); } private final class ToolCallAccumulator { diff --git a/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiProviderConfig.java b/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiProviderConfig.java index 4b65952c..116dcf2b 100644 --- a/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiProviderConfig.java +++ b/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/OpenAiProviderConfig.java @@ -33,4 +33,19 @@ public record OpenAiProviderConfig( websocketUrl = websocketUrl == null ? Optional.empty() : websocketUrl; compat = compat == null ? Map.of() : Map.copyOf(compat); } + + @Override + public String toString() { + return "OpenAiProviderConfig[provider=" + provider + + ", baseUrl=" + baseUrl + + ", websocketUrl=" + websocketUrl + + ", websocketPath=" + websocketPath + + ", apiKey=" + + ", requestStyle=" + requestStyle + + ", fallbackRequestStyle=" + fallbackRequestStyle + + ", transportMode=" + transportMode + + ", timeout=" + timeout + + ", maxRetries=" + maxRetries + + ", compat=]"; + } } diff --git a/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/TaggedThinkingContentParser.java b/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/TaggedThinkingContentParser.java new file mode 100644 index 00000000..42f94ae8 --- /dev/null +++ b/lypi-ai/src/main/java/cn/lypi/ai/provider/openai/TaggedThinkingContentParser.java @@ -0,0 +1,94 @@ +package cn.lypi.ai.provider.openai; + +import cn.lypi.contracts.model.AssistantStreamEvent; +import cn.lypi.contracts.model.TextDelta; +import cn.lypi.contracts.model.ThinkingDelta; +import java.util.ArrayList; +import java.util.List; + +public final class TaggedThinkingContentParser { + private static final String OPEN_TAG = ""; + private static final String CLOSE_TAG = ""; + + private final StringBuilder pending = new StringBuilder(); + private Mode mode = Mode.TEXT; + + public List accept(String content) { + if (content == null || content.isEmpty()) { + return List.of(); + } + pending.append(content); + List events = new ArrayList<>(); + while (!pending.isEmpty()) { + String expectedTag = mode == Mode.TEXT ? OPEN_TAG : CLOSE_TAG; + int tagIndex = pending.indexOf(expectedTag); + if (tagIndex >= 0) { + emit(events, pending.substring(0, tagIndex)); + pending.delete(0, tagIndex + expectedTag.length()); + mode = mode == Mode.TEXT ? Mode.THINKING : Mode.TEXT; + continue; + } + + int retainedLength = longestTagPrefixAtEnd(expectedTag); + int safeLength = pending.length() - retainedLength; + if (safeLength > 0) { + emit(events, pending.substring(0, safeLength)); + pending.delete(0, safeLength); + } + break; + } + return List.copyOf(events); + } + + public List finish() { + if (pending.isEmpty()) { + return List.of(); + } + List events = new ArrayList<>(); + emit(events, pending.toString()); + pending.setLength(0); + return List.copyOf(events); + } + + private int longestTagPrefixAtEnd(String tag) { + int maximum = Math.min(pending.length(), tag.length() - 1); + for (int length = maximum; length > 0; length--) { + int pendingStart = pending.length() - length; + boolean matches = true; + for (int index = 0; index < length; index++) { + if (pending.charAt(pendingStart + index) != tag.charAt(index)) { + matches = false; + break; + } + } + if (matches) { + return length; + } + } + return 0; + } + + private void emit(List events, String text) { + if (text.isEmpty()) { + return; + } + if (mode == Mode.TEXT) { + if (!events.isEmpty() && events.getLast() instanceof TextDelta previous) { + events.set(events.size() - 1, new TextDelta(previous.text() + text)); + } else { + events.add(new TextDelta(text)); + } + return; + } + if (!events.isEmpty() && events.getLast() instanceof ThinkingDelta previous) { + events.set(events.size() - 1, new ThinkingDelta(previous.text() + text)); + } else { + events.add(new ThinkingDelta(text)); + } + } + + private enum Mode { + TEXT, + THINKING + } +} diff --git a/lypi-ai/src/test/java/cn/lypi/ai/DefaultModelRegistryTest.java b/lypi-ai/src/test/java/cn/lypi/ai/DefaultModelRegistryTest.java index 2b3e19ee..8f8749e4 100644 --- a/lypi-ai/src/test/java/cn/lypi/ai/DefaultModelRegistryTest.java +++ b/lypi-ai/src/test/java/cn/lypi/ai/DefaultModelRegistryTest.java @@ -45,9 +45,10 @@ void findReturnsEmptyWhenSelectionIsUnavailable() { void registryExposesContractsModelCatalogPort() { ModelDescriptor descriptor = descriptor("openai", "gpt-5"); ModelCatalogPort catalog = new DefaultModelRegistry(List.of(descriptor)); + ModelSelection selection = new ModelSelection("openai", "gpt-5", ThinkingLevel.MEDIUM); - assertThat(catalog.find(new ModelSelection("openai", "gpt-5", ThinkingLevel.MEDIUM))) - .contains(descriptor); + assertThat(catalog.list()).containsExactly(descriptor); + assertThat(catalog.find(selection)).contains(descriptor); } @Test @@ -61,6 +62,31 @@ void registryDefensivelyCopiesDescriptors() { assertThat(registry.list()).containsExactly(descriptor); } + @Test + void replaceProviderRemovesStaleModelsAndPreservesOtherProviders() { + RuntimeModelRegistry registry = new DefaultModelRegistry(List.of(descriptor("fixed", "fixed-model"))); + + registry.replaceProvider("login-example", List.of( + descriptor("login-example", "model-a"), + descriptor("login-example", "model-b") + )); + registry.replaceProvider("login-example", List.of(descriptor("login-example", "model-c"))); + + assertThat(registry.list()) + .extracting(candidate -> candidate.provider() + "/" + candidate.modelId()) + .containsExactlyInAnyOrder("fixed/fixed-model", "login-example/model-c"); + assertThat(registry.find(new ModelSelection("login-example", "model-a", ThinkingLevel.OFF))).isEmpty(); + } + + @Test + void replaceProviderRejectsDescriptorsForAnotherProvider() { + RuntimeModelRegistry registry = new DefaultModelRegistry(List.of()); + + org.assertj.core.api.Assertions.assertThatThrownBy(() -> + registry.replaceProvider("login-example", List.of(descriptor("other", "model-a"))) + ).isInstanceOf(IllegalArgumentException.class); + } + private static ModelDescriptor descriptor(String provider, String modelId) { return new ModelDescriptor( provider, diff --git a/lypi-ai/src/test/java/cn/lypi/ai/ProviderAdapterApiProviderTest.java b/lypi-ai/src/test/java/cn/lypi/ai/ProviderAdapterApiProviderTest.java new file mode 100644 index 00000000..fa805395 --- /dev/null +++ b/lypi-ai/src/test/java/cn/lypi/ai/ProviderAdapterApiProviderTest.java @@ -0,0 +1,116 @@ +package cn.lypi.ai; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import cn.lypi.contracts.common.AbortSignal; +import cn.lypi.contracts.context.ContextBudget; +import cn.lypi.contracts.context.ContextSnapshot; +import cn.lypi.contracts.error.ModelProviderException; +import cn.lypi.contracts.model.ApiStyle; +import cn.lypi.contracts.model.AssistantEventStream; +import cn.lypi.contracts.model.AssistantStreamResult; +import cn.lypi.contracts.model.CostProfile; +import cn.lypi.contracts.model.ModelDescriptor; +import cn.lypi.contracts.model.ModelSelection; +import cn.lypi.contracts.model.ThinkingLevel; +import cn.lypi.contracts.prompt.SystemPrompt; +import cn.lypi.contracts.security.AgentMode; +import cn.lypi.contracts.security.PermissionMode; +import java.math.BigDecimal; +import java.net.URI; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import org.junit.jupiter.api.Test; + +class ProviderAdapterApiProviderTest { + @Test + void replacesAdapterForTheSameProvider() { + RecordingAdapter oldAdapter = new RecordingAdapter("login-example"); + RecordingAdapter newAdapter = new RecordingAdapter("login-example"); + ProviderAdapterApiProvider provider = new ProviderAdapterApiProvider( + ApiStyle.OPENAI_COMPATIBLE, + List.of(oldAdapter) + ); + + provider.replaceAdapter(newAdapter); + try (AssistantEventStream ignored = provider.stream(context(), descriptor(), () -> false)) { + assertThat(newAdapter.calls).isEqualTo(1); + assertThat(oldAdapter.calls).isZero(); + } + } + + @Test + void failsForAnAdapterThatHasNotBeenRegistered() { + ProviderAdapterApiProvider provider = new ProviderAdapterApiProvider( + ApiStyle.OPENAI_COMPATIBLE, + List.of() + ); + + assertThatThrownBy(() -> provider.stream(context(), descriptor(), () -> false)) + .isInstanceOfSatisfying(ModelProviderException.class, error -> + assertThat(error.errorId()).isEqualTo("provider.adapter_unavailable")); + } + + private static ModelDescriptor descriptor() { + return new ModelDescriptor( + "login-example", + "model-a", + URI.create("https://example.test/v1"), + ApiStyle.OPENAI_COMPATIBLE, + 0, + 0, + false, + false, + new CostProfile(BigDecimal.ZERO, BigDecimal.ZERO, "USD"), + Map.of() + ); + } + + private static ContextSnapshot context() { + return new ContextSnapshot( + new SystemPrompt("system", List.of(), "hash"), + List.of(), + new ModelSelection("login-example", "model-a", ThinkingLevel.OFF), + ThinkingLevel.OFF, + AgentMode.EXECUTE, + PermissionMode.ASK, + new ContextBudget(0, 1, 1, 1, 1, 0, 0, BigDecimal.ZERO) + ); + } + + private static final class RecordingAdapter implements ProviderAdapter { + private final String provider; + private int calls; + + private RecordingAdapter(String provider) { + this.provider = provider; + } + + @Override + public String provider() { + return provider; + } + + @Override + public AssistantEventStream stream(ContextSnapshot context, ModelDescriptor descriptor, AbortSignal signal) { + calls++; + return new AssistantEventStream() { + @Override + public java.util.Iterator iterator() { + return List.of().iterator(); + } + + @Override + public AssistantStreamResult result() { + return new AssistantStreamResult("", List.of(), Optional.empty(), Optional.empty(), true, false, Optional.empty()); + } + + @Override + public void close() { + } + }; + } + } +} diff --git a/lypi-ai/src/test/java/cn/lypi/ai/model/RemoteModelDescriptorSourceTest.java b/lypi-ai/src/test/java/cn/lypi/ai/model/RemoteModelDescriptorSourceTest.java index 73847547..31b5f55c 100644 --- a/lypi-ai/src/test/java/cn/lypi/ai/model/RemoteModelDescriptorSourceTest.java +++ b/lypi-ai/src/test/java/cn/lypi/ai/model/RemoteModelDescriptorSourceTest.java @@ -10,12 +10,23 @@ import java.time.Duration; import java.util.List; import java.util.Map; +import java.util.Optional; +import java.util.OptionalInt; import org.junit.jupiter.api.Test; class RemoteModelDescriptorSourceTest { @Test - void mapsDiscoveredModelIdsToDescriptors() { - RecordingDiscoveryClient client = new RecordingDiscoveryClient(List.of("remote-a", "remote-b")); + void mapsExplicitCapabilitiesAndDefaultsMissingMetadata() { + RecordingDiscoveryClient client = new RecordingDiscoveryClient(List.of( + new DiscoveredModel( + "remote-explicit", + OptionalInt.of(128_000), + OptionalInt.of(16_384), + Optional.of(false), + Optional.of(false) + ), + DiscoveredModel.idOnly("remote-defaulted") + )); RemoteModelDescriptorSource source = new RemoteModelDescriptorSource( true, "test-provider", @@ -28,12 +39,26 @@ void mapsDiscoveredModelIdsToDescriptors() { defaults() ); - assertThat(source.list()) + List descriptors = source.list(); + + assertThat(descriptors) .extracting(ModelDescriptor::provider, ModelDescriptor::modelId, ModelDescriptor::baseUrl, ModelDescriptor::apiStyle) .containsExactly( - org.assertj.core.groups.Tuple.tuple("test-provider", "remote-a", URI.create("https://provider.example/v1"), ApiStyle.OPENAI_COMPATIBLE), - org.assertj.core.groups.Tuple.tuple("test-provider", "remote-b", URI.create("https://provider.example/v1"), ApiStyle.OPENAI_COMPATIBLE) + org.assertj.core.groups.Tuple.tuple("test-provider", "remote-explicit", URI.create("https://provider.example/v1"), ApiStyle.OPENAI_COMPATIBLE), + org.assertj.core.groups.Tuple.tuple("test-provider", "remote-defaulted", URI.create("https://provider.example/v1"), ApiStyle.OPENAI_COMPATIBLE) ); + assertThat(descriptors.get(0)).satisfies(descriptor -> { + assertThat(descriptor.contextWindow()).isEqualTo(128_000); + assertThat(descriptor.maxOutputTokens()).isEqualTo(16_384); + assertThat(descriptor.supportsThinking()).isFalse(); + assertThat(descriptor.supportsImageInput()).isFalse(); + }); + assertThat(descriptors.get(1)).satisfies(descriptor -> { + assertThat(descriptor.contextWindow()).isEqualTo(256_000); + assertThat(descriptor.maxOutputTokens()).isEqualTo(8_192); + assertThat(descriptor.supportsThinking()).isTrue(); + assertThat(descriptor.supportsImageInput()).isTrue(); + }); assertThat(client.baseUrl).isEqualTo(URI.create("https://provider.example/v1")); assertThat(client.apiKey).isEqualTo("secret-key"); assertThat(client.paths).containsExactly("/models", "/model"); @@ -50,8 +75,8 @@ void descriptorsInheritDefaultsWithoutSecretCompat() { "secret-key", List.of("/models"), Duration.ofSeconds(3), - new RecordingDiscoveryClient(List.of("remote-a")), - new RemoteModelDescriptorSource.DescriptorDefaults( + new RecordingDiscoveryClient(List.of(DiscoveredModel.idOnly("remote-a"))), + new DiscoveredModelDefaults( 200_000, 32_000, true, @@ -81,7 +106,7 @@ void descriptorsInheritDefaultsWithoutSecretCompat() { @Test void returnsEmptyWhenDiscoveryDisabled() { - RecordingDiscoveryClient client = new RecordingDiscoveryClient(List.of("remote-a")); + RecordingDiscoveryClient client = new RecordingDiscoveryClient(List.of(DiscoveredModel.idOnly("remote-a"))); RemoteModelDescriptorSource source = new RemoteModelDescriptorSource( false, "test-provider", @@ -98,37 +123,37 @@ void returnsEmptyWhenDiscoveryDisabled() { assertThat(client.called).isFalse(); } - private static RemoteModelDescriptorSource.DescriptorDefaults defaults() { - return new RemoteModelDescriptorSource.DescriptorDefaults( - 128_000, - 16_384, - false, - false, + private static DiscoveredModelDefaults defaults() { + return new DiscoveredModelDefaults( + 256_000, + 8_192, + true, + true, new CostProfile(BigDecimal.ZERO, BigDecimal.ZERO, "USD"), Map.of() ); } private static final class RecordingDiscoveryClient extends RemoteModelDiscoveryClient { - private final List modelIds; + private final List models; private boolean called; private URI baseUrl; private String apiKey; private List paths; private Duration timeout; - private RecordingDiscoveryClient(List modelIds) { - this.modelIds = modelIds; + private RecordingDiscoveryClient(List models) { + this.models = models; } @Override - public List discover(URI baseUrl, String apiKey, List paths, Duration timeout) { + public List discoverModels(URI baseUrl, String apiKey, List paths, Duration timeout) { this.called = true; this.baseUrl = baseUrl; this.apiKey = apiKey; this.paths = paths; this.timeout = timeout; - return modelIds; + return models; } } } diff --git a/lypi-ai/src/test/java/cn/lypi/ai/model/RemoteModelDiscoveryClientTest.java b/lypi-ai/src/test/java/cn/lypi/ai/model/RemoteModelDiscoveryClientTest.java index 5f29be0f..f4fafd5d 100644 --- a/lypi-ai/src/test/java/cn/lypi/ai/model/RemoteModelDiscoveryClientTest.java +++ b/lypi-ai/src/test/java/cn/lypi/ai/model/RemoteModelDiscoveryClientTest.java @@ -1,7 +1,9 @@ package cn.lypi.ai.model; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import cn.lypi.contracts.error.ModelProviderException; import com.sun.net.httpserver.HttpExchange; import com.sun.net.httpserver.HttpServer; import java.io.IOException; @@ -10,6 +12,9 @@ import java.nio.charset.StandardCharsets; import java.time.Duration; import java.util.List; +import java.util.Optional; +import java.util.OptionalInt; +import java.util.concurrent.CopyOnWriteArrayList; import java.util.concurrent.atomic.AtomicReference; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Test; @@ -39,6 +44,197 @@ void discoversOpenAiDataModelsAndSendsAuthorizationHeader() throws IOException { assertThat(authorization.get()).isEqualTo("Bearer test-key"); } + @Test + void discoversOptionalCapabilitiesFromCommonModelMetadataShapes() throws IOException { + startServer(exchange -> respond(exchange, 200, """ + {"data":[ + { + "id":"direct-fields", + "context_window":131072, + "max_output_tokens":16384, + "supports_reasoning":false, + "supports_image_input":false + }, + { + "id":"nested-fields", + "context_length":262144, + "top_provider":{"max_completion_tokens":32768}, + "supported_parameters":["reasoning_effort"], + "architecture":{"input_modalities":["text","image"]} + }, + {"id":"id-only","object":"model"} + ]} + """)); + + RemoteModelDiscoveryClient client = new RemoteModelDiscoveryClient(); + + List models = client.discoverModels( + baseUrl(), + "test-key", + List.of("/models"), + Duration.ofSeconds(2) + ); + + assertThat(models).containsExactly( + new DiscoveredModel( + "direct-fields", + OptionalInt.of(131_072), + OptionalInt.of(16_384), + Optional.of(false), + Optional.of(false) + ), + new DiscoveredModel( + "nested-fields", + OptionalInt.of(262_144), + OptionalInt.of(32_768), + Optional.of(true), + Optional.of(true) + ), + DiscoveredModel.idOnly("id-only") + ); + } + + @Test + void readsTheFirstValidCapabilityFieldInDocumentedPriorityOrder() throws IOException { + startServer(exchange -> respond(exchange, 200, """ + {"data":[ + { + "id":"priority", + "contextWindow":-1, + "context_length":111, + "context_window":222, + "maxOutputTokens":"invalid", + "max_output_tokens":333, + "max_tokens":444, + "top_provider":{"max_completion_tokens":555}, + "supportsThinking":"invalid", + "supports_thinking":false, + "supportsReasoning":true, + "supports_reasoning":true, + "supportsImageInput":"invalid", + "supports_image_input":false, + "input_modalities":["text","image"] + }, + { + "id":"camel-case", + "contextWindow":64000, + "maxOutputTokens":4096, + "supportsThinking":true, + "supportsImageInput":true + }, + { + "id":"remaining-aliases", + "context_window":96000, + "max_tokens":2048, + "supportsReasoning":false, + "input_modalities":["text"] + } + ]} + """)); + + List models = new RemoteModelDiscoveryClient().discoverModels( + baseUrl(), + "test-key", + List.of("/models"), + Duration.ofSeconds(2) + ); + + assertThat(models).containsExactly( + new DiscoveredModel( + "priority", + OptionalInt.of(111), + OptionalInt.of(333), + Optional.of(false), + Optional.of(false) + ), + new DiscoveredModel( + "camel-case", + OptionalInt.of(64_000), + OptionalInt.of(4_096), + Optional.of(true), + Optional.of(true) + ), + new DiscoveredModel( + "remaining-aliases", + OptionalInt.of(96_000), + OptionalInt.of(2_048), + Optional.of(false), + Optional.of(false) + ) + ); + } + + @Test + void treatsInvalidOrAbsentCapabilitiesAsUnknown() throws IOException { + startServer(exchange -> respond(exchange, 200, """ + {"data":[ + { + "id":"invalid", + "contextWindow":0, + "context_length":-1, + "context_window":2147483648, + "maxOutputTokens":"8192", + "max_output_tokens":0, + "max_tokens":-2, + "top_provider":{"max_completion_tokens":9223372036854775808}, + "supportsThinking":"true", + "supported_parameters":["temperature"], + "supportsImageInput":1, + "input_modalities":["audio"] + }, + { + "id":"inferred", + "supported_parameters":["thinking"], + "architecture":{"input_modalities":["text","image"]} + }, + { + "id":"reasoning-alias", + "supported_parameters":["reasoning"] + } + ]} + """)); + + List models = new RemoteModelDiscoveryClient().discoverModels( + baseUrl(), + "test-key", + List.of("/models"), + Duration.ofSeconds(2) + ); + + assertThat(models.get(0)).isEqualTo(DiscoveredModel.idOnly("invalid")); + assertThat(models.get(1).supportsThinking()).contains(true); + assertThat(models.get(1).supportsImageInput()).contains(true); + assertThat(models.get(2).supportsThinking()).contains(true); + assertThat(models.get(2).supportsImageInput()).isEmpty(); + } + + @Test + void keepsTheFirstCompleteRecordForDuplicateIdsAndPreservesLegacyIdOrder() throws IOException { + startServer(exchange -> respond(exchange, 200, """ + {"data":[ + {"id":"model-a","context_window":128000,"supports_reasoning":true}, + {"id":"model-b"}, + {"id":"model-a","context_window":1,"supports_reasoning":false} + ]} + """)); + + RemoteModelDiscoveryClient client = new RemoteModelDiscoveryClient(); + + assertThat(client.discoverModels(baseUrl(), "test-key", List.of("/models"), Duration.ofSeconds(2))) + .containsExactly( + new DiscoveredModel( + "model-a", + OptionalInt.of(128_000), + OptionalInt.empty(), + Optional.of(true), + Optional.empty() + ), + DiscoveredModel.idOnly("model-b") + ); + assertThat(client.discover(baseUrl(), "test-key", List.of("/models"), Duration.ofSeconds(2))) + .containsExactly("model-a", "model-b"); + } + @Test void triesFallbackModelPathWhenModelsPathIsMissing() throws IOException { startServer(exchange -> { @@ -57,7 +253,7 @@ void triesFallbackModelPathWhenModelsPathIsMissing() throws IOException { @Test void parsesTopLevelStringArray() throws IOException { - startServer(exchange -> respond(exchange, 200, "[\"model-a\",\"model-b\"]")); + startServer(exchange -> respond(exchange, 200, "[\"model-a\",\"model-b\",\"model-a\"]")); RemoteModelDiscoveryClient client = new RemoteModelDiscoveryClient(); @@ -66,11 +262,82 @@ void parsesTopLevelStringArray() throws IOException { } @Test - void returnsEmptyListWhenNetworkFails() { + void throwsSanitizedErrorAfterEveryCandidatePathFails() throws IOException { + List requestedPaths = new CopyOnWriteArrayList<>(); + startServer(exchange -> { + String path = exchange.getRequestURI().getPath(); + requestedPaths.add(path); + respond(exchange, path.endsWith("/models") ? 401 : 404, ""); + }); + + assertThatThrownBy(() -> new RemoteModelDiscoveryClient().discover( + baseUrl(), + "secret-key", + List.of("/models", "/model"), + Duration.ofSeconds(2) + )) + .isInstanceOfSatisfying(ModelProviderException.class, error -> + assertThat(error.errorId()).isEqualTo("model.discovery_unavailable")) + .hasMessageContaining("/v1/models") + .hasMessageContaining("HTTP 401") + .hasMessageContaining("/v1/model") + .hasMessageContaining("HTTP 404") + .hasMessageNotContaining("secret-key"); + + assertThat(requestedPaths).containsExactly("/v1/models", "/v1/model"); + } + + @Test + void throwsWhenSuccessfulResponsesContainNoUsableModelIds() throws IOException { + startServer(exchange -> respond(exchange, 200, "{\"data\":[{\"object\":\"model\"}]}")); + + assertDiscoveryUnavailable( + () -> new RemoteModelDiscoveryClient().discover( + baseUrl(), + "test-key", + List.of("/models", "/model"), + Duration.ofSeconds(2) + ), + "no usable model ids" + ); + } + + @Test + void throwsWhenEveryCandidateReturnsInvalidJson() throws IOException { + startServer(exchange -> respond(exchange, 200, "{not-json")); + + assertDiscoveryUnavailable( + () -> new RemoteModelDiscoveryClient().discover( + baseUrl(), + "test-key", + List.of("/models", "/model"), + Duration.ofSeconds(2) + ), + "invalid JSON" + ); + } + + @Test + void throwsWhenNetworkFails() { RemoteModelDiscoveryClient client = new RemoteModelDiscoveryClient(); - assertThat(client.discover(URI.create("http://127.0.0.1:1/v1"), "test-key", List.of("/models"), Duration.ofMillis(100))) - .isEmpty(); + assertDiscoveryUnavailable( + () -> client.discover( + URI.create("http://127.0.0.1:1/v1"), + "test-key", + List.of("/models"), + Duration.ofMillis(100) + ), + "network error" + ); + } + + private static void assertDiscoveryUnavailable(ThrowingCall call, String diagnostic) { + assertThatThrownBy(call::run) + .isInstanceOfSatisfying(ModelProviderException.class, error -> + assertThat(error.errorId()).isEqualTo("model.discovery_unavailable")) + .hasMessageContaining(diagnostic) + .hasMessageNotContaining("test-key"); } private URI baseUrl() { @@ -95,4 +362,9 @@ private static void respond(HttpExchange exchange, int status, String body) thro private interface ExchangeHandler { void handle(HttpExchange exchange) throws IOException; } + + @FunctionalInterface + private interface ThrowingCall { + void run(); + } } diff --git a/lypi-ai/src/test/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsRequestBuilderTest.java b/lypi-ai/src/test/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsRequestBuilderTest.java index 61775c44..73334834 100644 --- a/lypi-ai/src/test/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsRequestBuilderTest.java +++ b/lypi-ai/src/test/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsRequestBuilderTest.java @@ -4,11 +4,14 @@ import cn.lypi.ai.provider.RequestStyle; import cn.lypi.ai.provider.TransportMode; +import cn.lypi.ai.spec.LypiAttachmentBlock; +import cn.lypi.ai.spec.LypiContentBlock; import cn.lypi.ai.spec.LypiGenerationOptions; import cn.lypi.ai.spec.LypiMessage; import cn.lypi.ai.spec.LypiModelRequest; import cn.lypi.ai.spec.LypiRole; import cn.lypi.ai.spec.LypiTextBlock; +import cn.lypi.ai.spec.LypiThinkingBlock; import cn.lypi.ai.spec.LypiToolCallBlock; import cn.lypi.ai.spec.LypiToolResultBlock; import cn.lypi.ai.spec.LypiToolSpec; @@ -23,6 +26,9 @@ import org.junit.jupiter.api.Test; class OpenAiChatCompletionsRequestBuilderTest { + private static final String REASONING_CONTENT_COMPAT = + "requires-reasoning-content-on-assistant-messages"; + @Test void buildsChatCompletionsRequestWithMessagesToolsAndReasoningEffort() { LypiToolSpec tool = new LypiToolSpec( @@ -137,6 +143,135 @@ void preservesAssistantContentWhenToolCallHistoryAlsoHasText() { assertThat(body.at("/messages/0/tool_calls/0/id").asText()).isEqualTo("call-1"); } + @Test + void replaysAssistantThinkingSeparatelyWhenCompatIsEnabled() { + LypiModelRequest request = assistantToolHistory(List.of( + new LypiThinkingBlock("private plan", Map.of()), + new LypiTextBlock("I will read it.", Map.of()), + new LypiToolCallBlock("call-1", "read_file", "", Map.of("input", Map.of("path", "pom.xml"))) + )); + + for (Object enabled : List.of(true, "true")) { + JsonNode body = new OpenAiChatCompletionsRequestBuilder().build( + request, + config(Map.of(REASONING_CONTENT_COMPAT, enabled)) + ); + + assertThat(body.get("messages")).hasSize(1); + assertThat(body.at("/messages/0/role").asText()).isEqualTo("assistant"); + assertThat(body.at("/messages/0/reasoning_content").asText()).isEqualTo("private plan"); + assertThat(body.at("/messages/0/content").asText()).isEqualTo("I will read it."); + assertThat(body.at("/messages/0/tool_calls/0/id").asText()).isEqualTo("call-1"); + assertThat(body.at("/messages/0/tool_calls/0/function/name").asText()).isEqualTo("read_file"); + assertThat(body.at("/messages/0/tool_calls/0/function/arguments").asText()) + .isEqualTo("{\"path\":\"pom.xml\"}"); + } + } + + @Test + void writesEmptyReasoningContentForAssistantToolCallWithoutThinking() { + LypiModelRequest request = assistantToolHistory(List.of( + new LypiToolCallBlock("call-1", "read_file", "", Map.of("input", Map.of("path", "pom.xml"))) + )); + + JsonNode body = new OpenAiChatCompletionsRequestBuilder().build( + request, + config(Map.of(REASONING_CONTENT_COMPAT, true)) + ); + + assertThat(body.at("/messages/0/reasoning_content").isTextual()).isTrue(); + assertThat(body.at("/messages/0/reasoning_content").asText()).isEmpty(); + assertThat(body.at("/messages/0/content").isNull()).isTrue(); + assertThat(body.at("/messages/0/tool_calls/0/id").asText()).isEqualTo("call-1"); + } + + @Test + void keepsThinkingInVisibleContentWhenReasoningCompatIsDisabled() { + LypiModelRequest request = assistantToolHistory(List.of( + new LypiThinkingBlock("private plan", Map.of()), + new LypiTextBlock("I will read it.", Map.of()), + new LypiToolCallBlock("call-1", "read_file", "", Map.of("input", Map.of("path", "pom.xml"))) + )); + + JsonNode body = new OpenAiChatCompletionsRequestBuilder().build(request, config()); + + assertThat(body.at("/messages/0/reasoning_content").isMissingNode()).isTrue(); + assertThat(body.at("/messages/0/content").asText()).isEqualTo("private plan\nI will read it."); + assertThat(body.at("/messages/0/tool_calls/0/id").asText()).isEqualTo("call-1"); + } + + @Test + void serializesUserTextAndImageAsOneMultimodalMessage() { + LypiModelRequest request = requestWithMessages(List.of(new LypiMessage( + LypiRole.USER, + List.of( + new LypiTextBlock("describe", Map.of()), + new LypiAttachmentBlock( + "image-1", + "attached image", + "image/png", + Map.of("imageUrl", "data:image/png;base64,AAA", "detail", "high") + ) + ), + Map.of() + ))); + + JsonNode body = new OpenAiChatCompletionsRequestBuilder().build(request, config()); + + assertThat(body.get("messages")).hasSize(1); + assertThat(body.at("/messages/0/role").asText()).isEqualTo("user"); + assertThat(body.at("/messages/0/content/0/type").asText()).isEqualTo("text"); + assertThat(body.at("/messages/0/content/0/text").asText()).isEqualTo("describe"); + assertThat(body.at("/messages/0/content/1/type").asText()).isEqualTo("image_url"); + assertThat(body.at("/messages/0/content/1/image_url/url").asText()) + .isEqualTo("data:image/png;base64,AAA"); + assertThat(body.at("/messages/0/content/1/image_url/detail").asText()).isEqualTo("high"); + } + + @Test + void preservesAttachmentTextWhenImageUrlIsMissing() { + LypiModelRequest request = requestWithMessages(List.of(new LypiMessage( + LypiRole.USER, + List.of(new LypiAttachmentBlock("attachment-1", "image description", "image/png", Map.of())), + Map.of() + ))); + + JsonNode body = new OpenAiChatCompletionsRequestBuilder().build(request, config()); + + assertThat(body.get("messages")).hasSize(1); + assertThat(body.at("/messages/0/content").isTextual()).isTrue(); + assertThat(body.at("/messages/0/content").asText()).isEqualTo("image description"); + } + + @Test + void sendsToolResultImagesAsAFollowingUserMessage() { + LypiModelRequest request = requestWithMessages(List.of(new LypiMessage( + LypiRole.TOOL_RESULT, + List.of( + new LypiToolResultBlock("call-1", "screenshot captured", false, Map.of()), + new LypiAttachmentBlock( + "image-1", + "captured screenshot", + "image/png", + Map.of("imageUrl", "data:image/png;base64,AAA") + ) + ), + Map.of() + ))); + + JsonNode body = new OpenAiChatCompletionsRequestBuilder().build(request, config()); + + assertThat(body.get("messages")).hasSize(2); + assertThat(body.at("/messages/0/role").asText()).isEqualTo("tool"); + assertThat(body.at("/messages/0/tool_call_id").asText()).isEqualTo("call-1"); + assertThat(body.at("/messages/0/content").asText()).isEqualTo("screenshot captured"); + assertThat(body.at("/messages/1/role").asText()).isEqualTo("user"); + assertThat(body.at("/messages/1/content/0/type").asText()).isEqualTo("image_url"); + assertThat(body.at("/messages/1/content/0/image_url/url").asText()) + .isEqualTo("data:image/png;base64,AAA"); + assertThat(body.at("/messages/1/content/0/image_url/detail").asText()).isEqualTo("high"); + } + @Test void omitsBlankSystemPromptAndReasoningWhenOff() { LypiModelRequest request = new LypiModelRequest( @@ -175,6 +310,10 @@ void includesPromptCacheKeyFromRequestMetadata() { } private static OpenAiProviderConfig config() { + return config(Map.of()); + } + + private static OpenAiProviderConfig config(Map compat) { return new OpenAiProviderConfig( "openai", URI.create("https://api.openai.com/v1"), @@ -186,6 +325,23 @@ private static OpenAiProviderConfig config() { TransportMode.AUTO, Duration.ofSeconds(30), 1, + compat + ); + } + + private static LypiModelRequest assistantToolHistory(List content) { + return requestWithMessages(List.of(new LypiMessage(LypiRole.ASSISTANT, content, Map.of()))); + } + + private static LypiModelRequest requestWithMessages(List messages) { + return new LypiModelRequest( + "req-thinking-history", + new ModelSelection("openai", "gpt-4o-mini", ThinkingLevel.HIGH), + ThinkingLevel.HIGH, + "", + messages, + List.of(), + LypiGenerationOptions.defaults(), Map.of() ); } diff --git a/lypi-ai/src/test/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsStreamNormalizerTest.java b/lypi-ai/src/test/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsStreamNormalizerTest.java index 3270d89c..f8618d2a 100644 --- a/lypi-ai/src/test/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsStreamNormalizerTest.java +++ b/lypi-ai/src/test/java/cn/lypi/ai/provider/openai/OpenAiChatCompletionsStreamNormalizerTest.java @@ -15,6 +15,86 @@ import org.junit.jupiter.api.Test; class OpenAiChatCompletionsStreamNormalizerTest { + @Test + void normalizesCompatibleReasoningShapesWithoutDuplicateDetails() { + OpenAiChatCompletionsStreamNormalizer normalizer = new OpenAiChatCompletionsStreamNormalizer(); + + List events = List.of( + normalizer.normalize(""" + {"choices":[{"index":0,"delta":{"reasoning":"think","reasoning_details":[{"type":"reasoning.text","text":"think"}]}}]} + """), + normalizer.normalize(""" + {"choices":[{"index":0,"delta":{"reasoning_text":"fallback"}}]} + """), + normalizer.normalize(""" + {"choices":[{"index":0,"delta":{"reasoning_details":[{"type":"reasoning.encrypted","data":"ignored"},{"type":"reasoning.text","text":"details-only"}]}}]} + """), + normalizer.normalize(""" + {"choices":[{"index":1,"delta":{"content":"answer"}}]} + """) + ).stream().flatMap(List::stream).toList(); + + assertThat(events).containsExactly( + new ThinkingDelta("think"), + new ThinkingDelta("fallback"), + new ThinkingDelta("details-only"), + new TextDelta("answer") + ); + } + + @Test + void separatesThinkingTagsAcrossContentChunkBoundaries() { + OpenAiChatCompletionsStreamNormalizer normalizer = new OpenAiChatCompletionsStreamNormalizer(); + + List events = List.of( + normalizer.normalize("{\"choices\":[{\"delta\":{\"content\":\"plananswer\"}}]}"), + normalizer.normalize("[DONE]") + ).stream().flatMap(List::stream).toList(); + + assertThat(events).containsExactly( + new ThinkingDelta("plan"), + new TextDelta("answer"), + new AssistantDone(Optional.empty(), Optional.of("stop")) + ); + } + + @Test + void flushesUnclosedTaggedContentInTheCurrentModeOnDone() { + OpenAiChatCompletionsStreamNormalizer normalizer = new OpenAiChatCompletionsStreamNormalizer(); + + List events = List.of( + normalizer.normalize("{\"choices\":[{\"delta\":{\"content\":\"beforeunfinished events = List.of( + normalizer.normalize("{\"choices\":[{\"delta\":{\"content\":\"planThe file is deliberately separate from user-maintained configuration so a login never rewrites it.

+ */ +public class LoginProviderPropertiesStore { + private static final String PROVIDERS_PREFIX = "lypi.ai.providers."; + private static final List DISCOVERY_PATHS = List.of("/models", "/model"); + private static final Set PRIVATE_FILE_PERMISSIONS = EnumSet.of( + PosixFilePermission.OWNER_READ, + PosixFilePermission.OWNER_WRITE + ); + + private final Path file; + + public LoginProviderPropertiesStore(Path userHome) { + this.file = Objects.requireNonNull(userHome, "userHome") + .resolve(".ly-pi") + .resolve("login-providers.properties"); + } + + public void save(String provider, URI baseUrl, String authKey) throws IOException { + String requiredProvider = requireProvider(provider); + URI requiredBaseUrl = Objects.requireNonNull(baseUrl, "baseUrl"); + String requiredAuthKey = Objects.requireNonNull(authKey, "authKey"); + Path directory = file.getParent(); + Path temporary = null; + boolean moved = false; + try { + Files.createDirectories(directory); + Properties properties = readProperties(); + String prefix = PROVIDERS_PREFIX + requiredProvider + "."; + removeProviderProperties(properties, prefix); + writeProviderProperties(properties, prefix, requiredBaseUrl, requiredAuthKey); + + temporary = Files.createTempFile(directory, ".login-providers-", ".tmp"); + setPrivatePermissionsIfSupported(temporary); + writeProperties(temporary, properties); + moveReplacing(temporary, file); + moved = true; + setPrivatePermissionsIfSupported(file); + } catch (IOException | RuntimeException ignored) { + throw persistenceFailure(); + } finally { + if (!moved && temporary != null) { + deleteQuietly(temporary); + } + } + } + + private Properties readProperties() throws IOException { + Properties properties = new Properties(); + if (!Files.exists(file)) { + return properties; + } + try (java.io.InputStream input = Files.newInputStream(file)) { + properties.load(input); + } + return properties; + } + + private static void removeProviderProperties(Properties properties, String prefix) { + for (String propertyName : List.copyOf(properties.stringPropertyNames())) { + if (propertyName.startsWith(prefix)) { + properties.remove(propertyName); + } + } + } + + private static void writeProviderProperties( + Properties properties, + String prefix, + URI baseUrl, + String authKey + ) { + properties.setProperty(prefix + "enabled", "true"); + properties.setProperty(prefix + "api-style", "openai_compatible"); + properties.setProperty(prefix + "base-url", baseUrl.toString()); + properties.setProperty(prefix + "api-key", authKey); + properties.setProperty(prefix + "request-style", "chat_completions"); + properties.setProperty(prefix + "fallback-request-style", "chat_completions"); + properties.setProperty(prefix + "transport", "sse"); + properties.setProperty(prefix + "model-discovery.enabled", "true"); + for (int index = 0; index < DISCOVERY_PATHS.size(); index++) { + properties.setProperty(prefix + "model-discovery.paths[" + index + "]", DISCOVERY_PATHS.get(index)); + } + properties.setProperty(prefix + "compat.requires-reasoning-content-on-assistant-messages", "true"); + } + + private static void writeProperties(Path temporary, Properties properties) throws IOException { + try (OutputStream output = Files.newOutputStream(temporary)) { + properties.store(output, null); + } + } + + private static void moveReplacing(Path temporary, Path target) throws IOException { + try { + Files.move(temporary, target, StandardCopyOption.ATOMIC_MOVE, StandardCopyOption.REPLACE_EXISTING); + } catch (AtomicMoveNotSupportedException error) { + Files.move(temporary, target, StandardCopyOption.REPLACE_EXISTING); + } + } + + private static void setPrivatePermissionsIfSupported(Path path) throws IOException { + if (Files.getFileAttributeView(path, PosixFileAttributeView.class) != null) { + Files.setPosixFilePermissions(path, PRIVATE_FILE_PERMISSIONS); + } + } + + private static void deleteQuietly(Path path) { + try { + Files.deleteIfExists(path); + } catch (IOException ignored) { + // The temporary path name is random and contains no provider credentials. + } + } + + private static String requireProvider(String provider) { + if (provider == null || provider.isBlank()) { + throw new IllegalArgumentException("provider is required"); + } + return provider; + } + + private static IOException persistenceFailure() { + return new IOException("Provider login properties could not be saved."); + } +} diff --git a/lypi-boot/src/main/java/cn/lypi/boot/ai/LyPiAiAutoConfiguration.java b/lypi-boot/src/main/java/cn/lypi/boot/ai/LyPiAiAutoConfiguration.java index 637ecfc2..a855a305 100644 --- a/lypi-boot/src/main/java/cn/lypi/boot/ai/LyPiAiAutoConfiguration.java +++ b/lypi-boot/src/main/java/cn/lypi/boot/ai/LyPiAiAutoConfiguration.java @@ -8,9 +8,11 @@ import cn.lypi.ai.ModelRegistry; import cn.lypi.ai.ProviderAdapter; import cn.lypi.ai.ProviderAdapterApiProvider; +import cn.lypi.ai.RuntimeModelRegistry; import cn.lypi.ai.model.BuiltinModelDescriptorSource; import cn.lypi.ai.model.CompatSanitizer; import cn.lypi.ai.model.CompositeModelDescriptorSource; +import cn.lypi.ai.model.DiscoveredModelDefaults; import cn.lypi.ai.model.ModelDescriptorSource; import cn.lypi.ai.model.RemoteModelDescriptorSource; import cn.lypi.ai.model.RemoteModelDiscoveryClient; @@ -33,14 +35,19 @@ import cn.lypi.contracts.model.ApiStyle; import cn.lypi.contracts.model.CostProfile; import cn.lypi.contracts.model.ModelDescriptor; +import cn.lypi.contracts.runtime.ProviderLoginPort; import java.math.BigDecimal; +import java.nio.file.Path; import java.time.Duration; import java.util.ArrayList; import java.util.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.Optional; +import java.util.Set; +import java.util.stream.Collectors; import org.springframework.beans.factory.annotation.Qualifier; +import org.springframework.beans.factory.ObjectProvider; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; import org.springframework.boot.context.properties.EnableConfigurationProperties; @@ -55,8 +62,8 @@ public class LyPiAiAutoConfiguration { private static final Duration BUILTIN_OPENAI_TIMEOUT = Duration.ofSeconds(30); @Bean - @ConditionalOnMissingBean - public ModelRegistry modelRegistry(LyPiAiProperties properties, RemoteModelDiscoveryClient discoveryClient) { + @ConditionalOnMissingBean(ModelRegistry.class) + public RuntimeModelRegistry modelRegistry(LyPiAiProperties properties, RemoteModelDiscoveryClient discoveryClient) { return new DefaultModelRegistry(modelDescriptorSource(properties, discoveryClient).list()); } @@ -72,22 +79,25 @@ public ModelPort modelPort( @Bean @ConditionalOnMissingBean public ApiProviderRegistry apiProviderRegistry( - @Qualifier("openAiCompatibleProviderAdapters") List openAiAdapters, + @Qualifier("openAiCompatibleApiProvider") ProviderAdapterApiProvider openAiCompatibleApiProvider, @Qualifier("anthropicProviderAdapters") List anthropicAdapters ) { List providers = new ArrayList<>(); - if (!openAiAdapters.isEmpty()) { - providers.add(new ProviderAdapterApiProvider(ApiStyle.OPENAI_COMPATIBLE, openAiAdapters)); - } + providers.add(openAiCompatibleApiProvider); if (!anthropicAdapters.isEmpty()) { providers.add(new ProviderAdapterApiProvider(ApiStyle.ANTHROPIC, anthropicAdapters)); } - if (providers.isEmpty()) { - return new DefaultApiProviderRegistry(List.of()); - } return new DefaultApiProviderRegistry(providers); } + @Bean(name = "openAiCompatibleApiProvider") + @ConditionalOnMissingBean(name = "openAiCompatibleApiProvider") + public ProviderAdapterApiProvider openAiCompatibleApiProvider( + @Qualifier("openAiCompatibleProviderAdapters") List openAiAdapters + ) { + return new ProviderAdapterApiProvider(ApiStyle.OPENAI_COMPATIBLE, openAiAdapters); + } + @Bean @ConditionalOnMissingBean(name = "openAiCompatibleProviderAdapters") public List openAiCompatibleProviderAdapters(LyPiAiProperties properties) { @@ -106,6 +116,35 @@ public RemoteModelDiscoveryClient remoteModelDiscoveryClient() { return new RemoteModelDiscoveryClient(); } + @Bean + @ConditionalOnMissingBean + public LoginProviderPropertiesStore loginProviderPropertiesStore() { + return new LoginProviderPropertiesStore(Path.of(System.getProperty("user.home", "."))); + } + + @Bean + @ConditionalOnMissingBean(ProviderLoginPort.class) + public ProviderLoginPort providerLoginPort( + LyPiAiProperties properties, + RemoteModelDiscoveryClient discoveryClient, + ObjectProvider modelRegistry, + @Qualifier("openAiCompatibleApiProvider") ObjectProvider openAiDispatcher, + LoginProviderPropertiesStore propertiesStore + ) { + RuntimeModelRegistry runtimeModelRegistry = modelRegistry.getIfAvailable(); + ProviderAdapterApiProvider dispatcher = openAiDispatcher.getIfAvailable(); + if (runtimeModelRegistry == null || dispatcher == null) { + return ProviderLoginPort.unavailable(); + } + return new OpenAiCompatibleProviderLoginService( + discoveryClient, + runtimeModelRegistry, + dispatcher, + propertiesStore, + descriptorDefaults(properties) + ); + } + @Bean @ConditionalOnMissingBean public CompactionSummarizer compactionSummarizer(ModelPort modelPort, LyPiAiProperties properties) { @@ -118,11 +157,49 @@ public CompactionSummarizer compactionSummarizer(ModelPort modelPort, LyPiAiProp } private ModelDescriptorSource modelDescriptorSource(LyPiAiProperties properties, RemoteModelDiscoveryClient discoveryClient) { - List sources = new ArrayList<>(); - sources.add(new StaticModelDescriptorSource(builtinModelDescriptors(properties))); - sources.add(new StaticModelDescriptorSource(remoteModelDescriptors(properties, discoveryClient))); - sources.add(new StaticModelDescriptorSource(modelDescriptors(properties))); - return new CompositeModelDescriptorSource(sources); + DiscoveredModelDefaults defaults = descriptorDefaults(properties); + List remote = remoteModelDescriptors(properties, discoveryClient, defaults); + Set discovered = remote.stream() + .map(model -> new ModelKey(model.provider(), model.modelId())) + .collect(Collectors.toUnmodifiableSet()); + List builtin = authoritativeLocalDescriptors( + properties, + builtinModelDescriptors(properties), + discovered + ); + List configured = authoritativeLocalDescriptors( + properties, + modelDescriptors(properties), + discovered + ); + return new CompositeModelDescriptorSource(List.of( + new StaticModelDescriptorSource(remote), + new StaticModelDescriptorSource(builtin), + new StaticModelDescriptorSource(configured) + )); + } + + private List authoritativeLocalDescriptors( + LyPiAiProperties properties, + List local, + Set discovered + ) { + Map providers = effectiveProviders(properties); + return local.stream() + .filter(descriptor -> { + ProviderProperties provider = providers.get(descriptor.provider()); + return !usesRemoteModelDiscovery(provider) + || discovered.contains(new ModelKey(descriptor.provider(), descriptor.modelId())); + }) + .toList(); + } + + private boolean usesRemoteModelDiscovery(ProviderProperties provider) { + return provider != null + && provider.isEnabled() + && provider.getBaseUrl() != null + && valueOrDefault(provider.getApiStyle(), ApiStyle.OPENAI_COMPATIBLE) == ApiStyle.OPENAI_COMPATIBLE + && provider.getModelDiscovery().isEnabled(); } private List builtinModelDescriptors(LyPiAiProperties properties) { @@ -156,7 +233,11 @@ private List modelDescriptors(LyPiAiProperties properties) { return descriptors; } - private List remoteModelDescriptors(LyPiAiProperties properties, RemoteModelDiscoveryClient discoveryClient) { + private List remoteModelDescriptors( + LyPiAiProperties properties, + RemoteModelDiscoveryClient discoveryClient, + DiscoveredModelDefaults defaults + ) { List descriptors = new ArrayList<>(); effectiveProviders(properties).forEach((providerName, provider) -> { if (!provider.isEnabled() || provider.getBaseUrl() == null || !provider.getModelDiscovery().isEnabled()) { @@ -165,7 +246,6 @@ private List remoteModelDescriptors(LyPiAiProperties properties if (valueOrDefault(provider.getApiStyle(), ApiStyle.OPENAI_COMPATIBLE) != ApiStyle.OPENAI_COMPATIBLE) { return; } - RemoteModelDescriptorSource.DescriptorDefaults defaults = descriptorDefaults(provider); descriptors.addAll(new RemoteModelDescriptorSource( true, providerName, @@ -181,19 +261,18 @@ private List remoteModelDescriptors(LyPiAiProperties properties return descriptors; } - private RemoteModelDescriptorSource.DescriptorDefaults descriptorDefaults(ProviderProperties provider) { - ModelProperties firstModel = provider.getModels().isEmpty() ? new ModelProperties() : provider.getModels().getFirst(); - return new RemoteModelDescriptorSource.DescriptorDefaults( - firstModel.getContextWindow(), - firstModel.getMaxOutputTokens(), - firstModel.isSupportsThinking(), - firstModel.isSupportsImageInput(), - new CostProfile( - valueOrDefault(firstModel.getInputTokenCost(), BigDecimal.ZERO), - valueOrDefault(firstModel.getOutputTokenCost(), BigDecimal.ZERO), - valueOrDefault(firstModel.getCurrency(), "USD") - ), - sanitizedCompat(provider.getCompat(), firstModel.getCompat()) + private DiscoveredModelDefaults descriptorDefaults(LyPiAiProperties properties) { + LyPiAiProperties.ModelDefaultsProperties defaults = properties.getModelDiscovery().getDefaults(); + if (defaults.getContextWindow() <= 0 || defaults.getMaxOutputTokens() <= 0) { + throw new IllegalArgumentException("Model discovery default token limits must be positive."); + } + return new DiscoveredModelDefaults( + defaults.getContextWindow(), + defaults.getMaxOutputTokens(), + defaults.isSupportsThinking(), + defaults.isSupportsImageInput(), + new CostProfile(BigDecimal.ZERO, BigDecimal.ZERO, "USD"), + Map.of() ); } @@ -403,4 +482,7 @@ private ModelDescriptor withProviderOverrides(ModelDescriptor descriptor, Provid private static T valueOrDefault(T value, T defaultValue) { return value == null ? defaultValue : value; } + + private record ModelKey(String provider, String modelId) { + } } diff --git a/lypi-boot/src/main/java/cn/lypi/boot/ai/LyPiAiProperties.java b/lypi-boot/src/main/java/cn/lypi/boot/ai/LyPiAiProperties.java index f0fdc7e8..89d626ac 100644 --- a/lypi-boot/src/main/java/cn/lypi/boot/ai/LyPiAiProperties.java +++ b/lypi-boot/src/main/java/cn/lypi/boot/ai/LyPiAiProperties.java @@ -2,6 +2,7 @@ import cn.lypi.ai.provider.RequestStyle; import cn.lypi.ai.provider.TransportMode; +import cn.lypi.ai.model.DiscoveredModelDefaults; import cn.lypi.agent.compact.CompactionSummaryFallbackPolicy; import cn.lypi.contracts.model.ApiStyle; import java.math.BigDecimal; @@ -17,6 +18,7 @@ public class LyPiAiProperties { private String defaultProvider; private String defaultModel; + private GlobalModelDiscoveryProperties modelDiscovery = new GlobalModelDiscoveryProperties(); private Map providers = new LinkedHashMap<>(); private CompactionSummaryProperties compactionSummary = new CompactionSummaryProperties(); @@ -36,6 +38,14 @@ public void setDefaultModel(String defaultModel) { this.defaultModel = defaultModel; } + public GlobalModelDiscoveryProperties getModelDiscovery() { + return modelDiscovery; + } + + public void setModelDiscovery(GlobalModelDiscoveryProperties modelDiscovery) { + this.modelDiscovery = modelDiscovery == null ? new GlobalModelDiscoveryProperties() : modelDiscovery; + } + public Map getProviders() { return providers; } @@ -66,6 +76,57 @@ public void setFallbackPolicy(CompactionSummaryFallbackPolicy fallbackPolicy) { } } + public static class GlobalModelDiscoveryProperties { + private ModelDefaultsProperties defaults = new ModelDefaultsProperties(); + + public ModelDefaultsProperties getDefaults() { + return defaults; + } + + public void setDefaults(ModelDefaultsProperties defaults) { + this.defaults = defaults == null ? new ModelDefaultsProperties() : defaults; + } + } + + public static class ModelDefaultsProperties { + private int contextWindow = DiscoveredModelDefaults.DEFAULT_CONTEXT_WINDOW; + private int maxOutputTokens = DiscoveredModelDefaults.DEFAULT_MAX_OUTPUT_TOKENS; + private boolean supportsThinking = DiscoveredModelDefaults.DEFAULT_SUPPORTS_THINKING; + private boolean supportsImageInput = DiscoveredModelDefaults.DEFAULT_SUPPORTS_IMAGE_INPUT; + + public int getContextWindow() { + return contextWindow; + } + + public void setContextWindow(int contextWindow) { + this.contextWindow = contextWindow; + } + + public int getMaxOutputTokens() { + return maxOutputTokens; + } + + public void setMaxOutputTokens(int maxOutputTokens) { + this.maxOutputTokens = maxOutputTokens; + } + + public boolean isSupportsThinking() { + return supportsThinking; + } + + public void setSupportsThinking(boolean supportsThinking) { + this.supportsThinking = supportsThinking; + } + + public boolean isSupportsImageInput() { + return supportsImageInput; + } + + public void setSupportsImageInput(boolean supportsImageInput) { + this.supportsImageInput = supportsImageInput; + } + } + public static class ProviderProperties { private boolean enabled; private boolean enabledConfigured; diff --git a/lypi-boot/src/main/java/cn/lypi/boot/ai/OpenAiCompatibleProviderLoginService.java b/lypi-boot/src/main/java/cn/lypi/boot/ai/OpenAiCompatibleProviderLoginService.java new file mode 100644 index 00000000..8eb59a69 --- /dev/null +++ b/lypi-boot/src/main/java/cn/lypi/boot/ai/OpenAiCompatibleProviderLoginService.java @@ -0,0 +1,222 @@ +package cn.lypi.boot.ai; + +import cn.lypi.ai.ProviderAdapterApiProvider; +import cn.lypi.ai.RuntimeModelRegistry; +import cn.lypi.ai.model.DiscoveredModel; +import cn.lypi.ai.model.DiscoveredModelDefaults; +import cn.lypi.ai.model.DiscoveredModelDescriptorMapper; +import cn.lypi.ai.model.RemoteModelDiscoveryClient; +import cn.lypi.ai.provider.RequestStyle; +import cn.lypi.ai.provider.TransportMode; +import cn.lypi.ai.provider.openai.OpenAiCompatibleProviderAdapter; +import cn.lypi.ai.provider.openai.OpenAiProviderConfig; +import cn.lypi.ai.transport.HttpSseProviderTransport; +import cn.lypi.ai.transport.WebSocketProviderTransport; +import cn.lypi.contracts.error.ErrorSeverity; +import cn.lypi.contracts.error.ModelProviderException; +import cn.lypi.contracts.model.ApiStyle; +import cn.lypi.contracts.model.ModelDescriptor; +import cn.lypi.contracts.runtime.ProviderLoginPort; +import cn.lypi.contracts.runtime.ProviderLoginResult; +import java.io.IOException; +import java.net.URI; +import java.net.URISyntaxException; +import java.time.Duration; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.Objects; +import java.util.Optional; +import java.util.regex.Pattern; + +/** Registers a verified OpenAI-compatible Chat Completions provider for the current process. */ +public final class OpenAiCompatibleProviderLoginService implements ProviderLoginPort { + private static final Pattern CHANNEL_NAME = Pattern.compile("[a-z0-9][a-z0-9_-]{0,63}"); + private static final List DISCOVERY_PATHS = List.of("/models", "/model"); + private static final Duration REQUEST_TIMEOUT = Duration.ofSeconds(30); + private static final int MAX_RETRIES = 3; + + private final RemoteModelDiscoveryClient discoveryClient; + private final RuntimeModelRegistry modelRegistry; + private final ProviderAdapterApiProvider openAiDispatcher; + private final LoginProviderPropertiesStore propertiesStore; + private final DiscoveredModelDefaults defaults; + + public OpenAiCompatibleProviderLoginService( + RemoteModelDiscoveryClient discoveryClient, + RuntimeModelRegistry modelRegistry, + ProviderAdapterApiProvider openAiDispatcher, + LoginProviderPropertiesStore propertiesStore, + DiscoveredModelDefaults defaults + ) { + this.discoveryClient = Objects.requireNonNull(discoveryClient, "discoveryClient"); + this.modelRegistry = Objects.requireNonNull(modelRegistry, "modelRegistry"); + this.openAiDispatcher = Objects.requireNonNull(openAiDispatcher, "openAiDispatcher"); + this.propertiesStore = Objects.requireNonNull(propertiesStore, "propertiesStore"); + this.defaults = Objects.requireNonNull(defaults, "defaults"); + } + + @Override + public ProviderLoginResult register(String channelName, String rawBaseUrl, String authKey) { + String provider = requireChannelName(channelName); + URI baseUrl = normalizeBaseUrl(rawBaseUrl); + String requiredAuthKey = requireAuthKey(authKey); + List discovered = discoverModels(baseUrl, requiredAuthKey); + DiscoveredModelDescriptorMapper mapper = new DiscoveredModelDescriptorMapper( + provider, + baseUrl, + ApiStyle.OPENAI_COMPATIBLE, + defaults + ); + List descriptors = discovered.stream().map(mapper::map).toList(); + OpenAiCompatibleProviderAdapter adapter = chatCompletionsAdapter(provider, baseUrl, requiredAuthKey); + + try { + propertiesStore.save(provider, baseUrl, requiredAuthKey); + } catch (IOException | RuntimeException error) { + throw providerLoginFailure( + "provider.login_persistence_failed", + "Provider login could not be saved." + ); + } + + openAiDispatcher.replaceAdapter(adapter); + modelRegistry.replaceProvider(provider, descriptors); + return new ProviderLoginResult(provider, descriptors); + } + + private List discoverModels(URI baseUrl, String authKey) { + List discovered; + try { + discovered = discoveryClient.discoverModels(baseUrl, authKey, DISCOVERY_PATHS, REQUEST_TIMEOUT); + } catch (ModelProviderException error) { + if (error.getMessage() != null && error.getMessage().contains(authKey)) { + throw providerLoginFailure( + "provider.login_discovery_failed", + "Provider model discovery failed." + ); + } + throw error; + } catch (RuntimeException error) { + throw providerLoginFailure( + "provider.login_discovery_failed", + "Provider model discovery failed." + ); + } + if (discovered.isEmpty()) { + throw providerLoginFailure( + "model.discovery_unavailable", + "Remote model discovery returned no usable models." + ); + } + return List.copyOf(discovered); + } + + private static OpenAiCompatibleProviderAdapter chatCompletionsAdapter( + String provider, + URI baseUrl, + String authKey + ) { + OpenAiProviderConfig config = new OpenAiProviderConfig( + provider, + baseUrl, + Optional.empty(), + "/v1/responses", + authKey, + RequestStyle.CHAT_COMPLETIONS, + RequestStyle.CHAT_COMPLETIONS, + TransportMode.SSE, + REQUEST_TIMEOUT, + MAX_RETRIES, + Map.of("requires-reasoning-content-on-assistant-messages", true) + ); + return new OpenAiCompatibleProviderAdapter( + config, + new WebSocketProviderTransport(), + new HttpSseProviderTransport(), + new HttpSseProviderTransport() + ); + } + + private static String requireChannelName(String channelName) { + if (channelName == null || !CHANNEL_NAME.matcher(channelName).matches()) { + throw providerLoginFailure( + "provider.login_invalid_channel_name", + "Provider channel name must match [a-z0-9][a-z0-9_-]{0,63}." + ); + } + return channelName; + } + + private static URI normalizeBaseUrl(String rawBaseUrl) { + if (rawBaseUrl == null || rawBaseUrl.isBlank()) { + throw providerLoginFailure( + "provider.login_invalid_base_url", + "Provider base URL must be an absolute HTTP(S) URL." + ); + } + URI parsed; + try { + parsed = URI.create(rawBaseUrl.trim()); + } catch (IllegalArgumentException error) { + throw providerLoginFailure( + "provider.login_invalid_base_url", + "Provider base URL must be an absolute HTTP(S) URL." + ); + } + if (!parsed.isAbsolute() + || parsed.isOpaque() + || parsed.getHost() == null + || parsed.getRawUserInfo() != null + || parsed.getRawQuery() != null + || parsed.getRawFragment() != null + || !("http".equalsIgnoreCase(parsed.getScheme()) || "https".equalsIgnoreCase(parsed.getScheme()))) { + throw providerLoginFailure( + "provider.login_invalid_base_url", + "Provider base URL must be an absolute HTTP(S) URL." + ); + } + String path = trimTrailingSlashes(parsed.getPath()); + try { + return new URI( + parsed.getScheme().toLowerCase(Locale.ROOT), + null, + parsed.getHost().toLowerCase(Locale.ROOT), + parsed.getPort(), + path, + null, + null + ); + } catch (URISyntaxException error) { + throw providerLoginFailure( + "provider.login_invalid_base_url", + "Provider base URL must be an absolute HTTP(S) URL." + ); + } + } + + private static String trimTrailingSlashes(String path) { + if (path == null || path.isEmpty() || "/".equals(path)) { + return null; + } + int end = path.length(); + while (end > 0 && path.charAt(end - 1) == '/') { + end--; + } + return end == 0 ? null : path.substring(0, end); + } + + private static String requireAuthKey(String authKey) { + if (authKey == null || authKey.isBlank()) { + throw providerLoginFailure( + "provider.login_invalid_auth_key", + "Provider auth key is required." + ); + } + return authKey; + } + + private static ModelProviderException providerLoginFailure(String errorId, String message) { + return new ModelProviderException(errorId, ErrorSeverity.ERROR, false, message); + } +} diff --git a/lypi-boot/src/main/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfiguration.java b/lypi-boot/src/main/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfiguration.java index 62698d8c..91bbdf99 100644 --- a/lypi-boot/src/main/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfiguration.java +++ b/lypi-boot/src/main/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfiguration.java @@ -19,6 +19,7 @@ import cn.lypi.contracts.runtime.CompactStateBackfillPort; import cn.lypi.contracts.runtime.CompactionRuntimePort; import cn.lypi.contracts.runtime.LyPiRuntime; +import cn.lypi.contracts.runtime.ProviderLoginPort; import cn.lypi.contracts.runtime.ResourceRuntimePort; import cn.lypi.contracts.runtime.SecurityRuntimePort; import cn.lypi.contracts.runtime.SessionManagerFactoryPort; @@ -429,9 +430,17 @@ public AppEntry appEntry( public JLineTuiTransportFactory jLineTuiTransportFactory( SessionManagerPort sessionManager, ResourceRuntimePort resourceRuntime, - CompactionRuntimePort compactionRuntime + CompactionRuntimePort compactionRuntime, + ObjectProvider modelCatalog, + ObjectProvider providerLogin ) { - return RuntimeBeanFactories.jLineTuiTransportFactory(sessionManager, resourceRuntime, compactionRuntime); + return RuntimeBeanFactories.jLineTuiTransportFactory( + sessionManager, + resourceRuntime, + compactionRuntime, + modelCatalog.getIfAvailable(), + providerLogin.getIfAvailable() + ); } /** diff --git a/lypi-boot/src/main/java/cn/lypi/boot/runtime/RuntimeBeanFactories.java b/lypi-boot/src/main/java/cn/lypi/boot/runtime/RuntimeBeanFactories.java index 8b29b864..3baae96c 100644 --- a/lypi-boot/src/main/java/cn/lypi/boot/runtime/RuntimeBeanFactories.java +++ b/lypi-boot/src/main/java/cn/lypi/boot/runtime/RuntimeBeanFactories.java @@ -33,6 +33,7 @@ import cn.lypi.contracts.runtime.CompactStateBackfillPort; import cn.lypi.contracts.runtime.CompactionRuntimePort; import cn.lypi.contracts.runtime.LyPiRuntime; +import cn.lypi.contracts.runtime.ProviderLoginPort; import cn.lypi.contracts.runtime.ResourceRuntimePort; import cn.lypi.contracts.runtime.SecurityRuntimePort; import cn.lypi.contracts.runtime.SessionManagerFactoryPort; @@ -525,7 +526,9 @@ static AppEntry appEntry( static JLineTuiTransportFactory jLineTuiTransportFactory( SessionManagerPort sessionManager, ResourceRuntimePort resourceRuntime, - CompactionRuntimePort compactionRuntime + CompactionRuntimePort compactionRuntime, + ModelCatalogPort modelCatalog, + ProviderLoginPort providerLogin ) { return (state, core, events, terminal, diffViewProvider, resumeController, newSessionController, slashCommands) -> JLineTuiTransport.open( @@ -539,7 +542,9 @@ static JLineTuiTransportFactory jLineTuiTransportFactory( newSessionController, sessionManager, resourceRuntime, - compactionRuntime + compactionRuntime, + modelCatalog, + providerLogin ); } diff --git a/lypi-boot/src/main/resources/application.yml b/lypi-boot/src/main/resources/application.yml index 90fb7556..260a5458 100644 --- a/lypi-boot/src/main/resources/application.yml +++ b/lypi-boot/src/main/resources/application.yml @@ -1,6 +1,8 @@ spring: config: - import: optional:file:${user.home}/.ly-pi/application.yml + import: + - optional:file:${user.home}/.ly-pi/login-providers.properties + - optional:file:${user.home}/.ly-pi/application.yml main: banner-mode: off diff --git a/lypi-boot/src/main/resources/application.yml.example b/lypi-boot/src/main/resources/application.yml.example index 710883d8..39a11b18 100644 --- a/lypi-boot/src/main/resources/application.yml.example +++ b/lypi-boot/src/main/resources/application.yml.example @@ -119,6 +119,15 @@ # # NOTE: 确定性摘要器已删除,fallback_deterministic 仅保留旧配置兼容;当前行为与 skip_compaction 一样回到原上下文。 # fallback-policy: fallback_deterministic # +# model-discovery: +# # 重要性:可选覆盖。远端模型目录缺少能力字段时使用这组全局默认描述。 +# # 优先级:静态同名完整模型描述 > 远端显式字段 > 本组默认值。 +# defaults: +# context-window: 256000 +# max-output-tokens: 8192 +# supports-thinking: true +# supports-image-input: true +# # # 重要性:可选覆盖或扩展。providers 用于覆盖内置 Provider 或新增 OpenAI 兼容或 Anthropic Provider。 # # 说明:内置 openai 可省略;配置 openai 会覆盖同名内置 adapter,配置新名称会追加 provider。 # # 关闭内置 openai 时,只需要取消注释: @@ -225,7 +234,7 @@ # api-style: openai_compatible # request-style: chat_completions # fallback-request-style: chat_completions -# transport: auto +# transport: sse # base-url: https://api.fixture.example/v1 # websocket-path: /v1/responses # websocket-url: @@ -234,10 +243,16 @@ # max-retries: 1 # compat: # vendor: fixture +# # 部分 Chat Completions 服务要求 assistant 工具调用历史携带 reasoning_content。 +# requires-reasoning-content-on-assistant-messages: true # model-discovery: -# enabled: false +# # 启动时按顺序请求候选路径;所有路径都没有返回有效模型时,应用启动失败。 +# enabled: true # paths: # - /models +# - /model +# # 开启 model-discovery 后,静态 models 作为远端同名模型的完整描述覆盖; +# # 远端未返回的静态 model-id 仍不会进入模型目录。 # models: # - model-id: fixture-model # context-window: 64000 diff --git a/lypi-boot/src/test/java/cn/lypi/boot/ApplicationExampleConfigTest.java b/lypi-boot/src/test/java/cn/lypi/boot/ApplicationExampleConfigTest.java index c77106f0..2e3ac994 100644 --- a/lypi-boot/src/test/java/cn/lypi/boot/ApplicationExampleConfigTest.java +++ b/lypi-boot/src/test/java/cn/lypi/boot/ApplicationExampleConfigTest.java @@ -19,11 +19,13 @@ import java.nio.charset.StandardCharsets; import java.time.Duration; import java.util.Map; +import java.util.stream.Collectors; import org.junit.jupiter.api.Test; import org.springframework.boot.context.properties.bind.Binder; import org.springframework.boot.context.properties.source.MapConfigurationPropertySource; import org.springframework.boot.env.YamlPropertySourceLoader; import org.springframework.core.env.StandardEnvironment; +import org.springframework.core.io.ByteArrayResource; import org.springframework.core.io.ClassPathResource; class ApplicationExampleConfigTest { @@ -109,6 +111,58 @@ void applicationExampleKeepsOpenAiThinkingSupportSeparateFromAnthropicLimitation assertThat(openAiBlock).doesNotContain("Anthropic extended thinking"); } + @Test + void applicationExampleDocumentsDiscoveredChatCompletionsProvider() throws IOException { + String example = new ClassPathResource("application.yml.example").getContentAsString(StandardCharsets.UTF_8); + String fixtureBlock = example.substring( + example.indexOf("# fixture:"), + example.indexOf("# anthropic:") + ); + + assertThat(fixtureBlock).contains("# request-style: chat_completions"); + assertThat(fixtureBlock).contains("# fallback-request-style: chat_completions"); + assertThat(fixtureBlock).contains("# transport: sse"); + assertThat(fixtureBlock).contains("# enabled: true"); + assertThat(fixtureBlock).contains("# - /models"); + assertThat(fixtureBlock).contains("# - /model"); + assertThat(fixtureBlock).contains("作为远端同名模型的完整描述覆盖"); + assertThat(fixtureBlock).contains("远端未返回的静态 model-id 仍不会进入模型目录"); + assertThat(fixtureBlock) + .contains("# requires-reasoning-content-on-assistant-messages: true"); + } + + @Test + void applicationExampleParsesDiscoveryDefaultsAndChatCompatibility() throws IOException { + StandardEnvironment environment = environmentForAiExample(); + + assertThat(environment.getProperty( + "lypi.ai.model-discovery.defaults.context-window", + Integer.class + )).isEqualTo(256000); + assertThat(environment.getProperty( + "lypi.ai.model-discovery.defaults.max-output-tokens", + Integer.class + )).isEqualTo(8192); + assertThat(environment.getProperty( + "lypi.ai.model-discovery.defaults.supports-thinking", + Boolean.class + )).isTrue(); + assertThat(environment.getProperty( + "lypi.ai.model-discovery.defaults.supports-image-input", + Boolean.class + )).isTrue(); + assertThat(environment.getProperty("lypi.ai.providers.fixture.request-style")) + .isEqualTo("chat_completions"); + assertThat(environment.getProperty("lypi.ai.providers.fixture.fallback-request-style")) + .isEqualTo("chat_completions"); + assertThat(environment.getProperty("lypi.ai.providers.fixture.transport")) + .isEqualTo("sse"); + assertThat(environment.getProperty( + "lypi.ai.providers.fixture.compat.requires-reasoning-content-on-assistant-messages", + Boolean.class + )).isTrue(); + } + @Test void applicationExampleDocumentsPermissionsAtLypiTopLevel() throws IOException { String example = new ClassPathResource("application.yml.example").getContentAsString(StandardCharsets.UTF_8); @@ -263,4 +317,24 @@ private Binder binderForExample() throws IOException { .forEach(environment.getPropertySources()::addLast); return Binder.get(environment); } + + private StandardEnvironment environmentForAiExample() throws IOException { + String example = new ClassPathResource("application.yml.example") + .getContentAsString(StandardCharsets.UTF_8); + String aiBlock = example.substring( + example.indexOf("# ai:"), + example.indexOf("# tool:") + ); + String yaml = "lypi:\n" + aiBlock.lines() + .map(line -> line.equals("#") ? "" : line.startsWith("# ") ? line.substring(2) : line) + .collect(Collectors.joining("\n")); + StandardEnvironment environment = new StandardEnvironment(); + new YamlPropertySourceLoader() + .load( + "application-example-ai", + new ByteArrayResource(yaml.getBytes(StandardCharsets.UTF_8)) + ) + .forEach(environment.getPropertySources()::addLast); + return environment; + } } diff --git a/lypi-boot/src/test/java/cn/lypi/boot/UserRootConfigurationTest.java b/lypi-boot/src/test/java/cn/lypi/boot/UserRootConfigurationTest.java index b1bd5e3a..0f6e9a11 100644 --- a/lypi-boot/src/test/java/cn/lypi/boot/UserRootConfigurationTest.java +++ b/lypi-boot/src/test/java/cn/lypi/boot/UserRootConfigurationTest.java @@ -82,6 +82,33 @@ void systemPropertyOverridesUserRootConfiguration() throws Exception { .isEqualTo(CompactionSummaryFallbackPolicy.FALLBACK_DETERMINISTIC)); } + @Test + void importsManagedLoginProviderConfigurationBeforeUserConfiguration() throws Exception { + Path home = Files.createDirectories(tempDir.resolve("login-provider-home")); + Path configRoot = Files.createDirectories(home.resolve(".ly-pi")); + Files.writeString(configRoot.resolve("login-providers.properties"), """ + lypi.ai.providers.login-fixture.enabled=true + lypi.ai.providers.login-fixture.base-url=https://generated.test/v1 + """); + Files.writeString(configRoot.resolve("application.yml"), """ + lypi: + ai: + providers: + login-fixture: + base-url: https://user.test/v1 + """); + + runner(home).run(context -> { + LyPiAiProperties.ProviderProperties provider = context.getBean(LyPiAiProperties.class) + .getProviders() + .get("login-fixture"); + + assertThat(provider).isNotNull(); + assertThat(provider.isEnabled()).isTrue(); + assertThat(provider.getBaseUrl()).hasToString("https://user.test/v1"); + }); + } + private ApplicationContextRunner runner(Path home) { return new ApplicationContextRunner() .withInitializer(new ConfigDataApplicationContextInitializer()) diff --git a/lypi-boot/src/test/java/cn/lypi/boot/ai/LyPiAiAutoConfigurationTest.java b/lypi-boot/src/test/java/cn/lypi/boot/ai/LyPiAiAutoConfigurationTest.java index 8d505fdd..bb95f7e3 100644 --- a/lypi-boot/src/test/java/cn/lypi/boot/ai/LyPiAiAutoConfigurationTest.java +++ b/lypi-boot/src/test/java/cn/lypi/boot/ai/LyPiAiAutoConfigurationTest.java @@ -1,12 +1,17 @@ package cn.lypi.boot.ai; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; import cn.lypi.ai.ApiProviderRegistry; import cn.lypi.ai.ModelPort; import cn.lypi.ai.ModelRegistry; +import cn.lypi.ai.ProviderAdapterApiProvider; +import cn.lypi.ai.RuntimeModelRegistry; +import cn.lypi.ai.model.DiscoveredModel; import cn.lypi.ai.model.RemoteModelDiscoveryClient; import cn.lypi.ai.provider.RequestStyle; +import cn.lypi.ai.provider.TransportMode; import cn.lypi.ai.provider.anthropic.AnthropicCompatibleProviderAdapter; import cn.lypi.ai.provider.anthropic.AnthropicProviderConfig; import cn.lypi.ai.provider.openai.OpenAiCompatibleProviderAdapter; @@ -14,16 +19,29 @@ import cn.lypi.agent.compact.AiCompactionSummarizer; import cn.lypi.agent.compact.CompactionSummarizer; import cn.lypi.agent.compact.CompactionSummaryFallbackPolicy; +import cn.lypi.contracts.error.ErrorSeverity; +import cn.lypi.contracts.error.ModelProviderException; import cn.lypi.contracts.model.ModelDescriptor; +import cn.lypi.contracts.runtime.ProviderLoginPort; +import cn.lypi.contracts.runtime.ProviderLoginResult; import java.lang.reflect.Field; import java.net.URI; +import java.nio.file.Files; +import java.nio.file.Path; import java.time.Duration; import java.util.List; +import java.util.Optional; +import java.util.OptionalInt; +import java.util.concurrent.atomic.AtomicInteger; import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; import org.springframework.boot.test.context.ConfigDataApplicationContextInitializer; import org.springframework.boot.test.context.runner.ApplicationContextRunner; class LyPiAiAutoConfigurationTest { + @TempDir + Path tempDir; + private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() .withUserConfiguration(LyPiAiAutoConfiguration.class) .withPropertyValues( @@ -160,6 +178,87 @@ void explicitOpenAiProviderDisableRemovesBuiltInAdapter() { }); } + @Test + void exposesMutableOpenAiDispatcherEvenWithoutInitialOpenAiAdapter() { + new ApplicationContextRunner() + .withUserConfiguration(LyPiAiAutoConfiguration.class) + .withPropertyValues("lypi.ai.providers.openai.enabled=false") + .run(context -> { + assertThat(context).hasSingleBean(ProviderAdapterApiProvider.class); + assertThat(context).hasSingleBean(RuntimeModelRegistry.class); + assertThat(context).hasSingleBean(ProviderLoginPort.class); + assertThat(context.getBean("openAiCompatibleProviderAdapters", List.class)).isEmpty(); + assertThat(context.getBean(ApiProviderRegistry.class) + .find(cn.lypi.contracts.model.ApiStyle.OPENAI_COMPATIBLE)).isPresent(); + }); + } + + @Test + void exposesUnavailableLoginPortWhenModelRegistryIsReplacedWithoutRuntimeMutationSupport() { + new ApplicationContextRunner() + .withUserConfiguration(LyPiAiAutoConfiguration.class) + .withBean(ModelRegistry.class, () -> new ModelRegistry() { + @Override + public List list() { + return List.of(); + } + + @Override + public java.util.Optional find(cn.lypi.contracts.model.ModelSelection selection) { + return java.util.Optional.empty(); + } + }) + .run(context -> assertThatThrownBy(() -> context.getBean(ProviderLoginPort.class) + .register("zen", "https://example.test/v1", "fixture-key")) + .isInstanceOf(IllegalStateException.class) + .hasMessage("provider login is unavailable")); + } + + @Test + void importsManagedLoginProviderAndDiscoversItsModelsAtStartup() throws Exception { + Path home = Files.createDirectories(tempDir.resolve("home")); + Path configRoot = Files.createDirectories(home.resolve(".ly-pi")); + Files.writeString(configRoot.resolve("login-providers.properties"), """ + lypi.ai.providers.login-fixture.enabled=true + lypi.ai.providers.login-fixture.api-style=openai_compatible + lypi.ai.providers.login-fixture.base-url=https://fixture.test/v1 + lypi.ai.providers.login-fixture.api-key=fixture-login-key + lypi.ai.providers.login-fixture.request-style=chat_completions + lypi.ai.providers.login-fixture.fallback-request-style=chat_completions + lypi.ai.providers.login-fixture.transport=sse + lypi.ai.providers.login-fixture.model-discovery.enabled=true + lypi.ai.providers.login-fixture.model-discovery.paths[0]=/models + lypi.ai.providers.login-fixture.model-discovery.paths[1]=/model + lypi.ai.providers.login-fixture.compat.requires-reasoning-content-on-assistant-messages=true + """); + + new ApplicationContextRunner() + .withInitializer(new ConfigDataApplicationContextInitializer()) + .withUserConfiguration(LyPiAiAutoConfiguration.class) + .withBean(RemoteModelDiscoveryClient.class, () -> new FixedRemoteModelDiscoveryClient("discovered-login-model")) + .withSystemProperties("user.home=" + home) + .run(context -> { + ModelDescriptor descriptor = model( + context.getBean(ModelRegistry.class), + "login-fixture", + "discovered-login-model" + ); + + assertThat(descriptor.baseUrl()).hasToString("https://fixture.test/v1"); + assertThat(context.getBean(ProviderLoginPort.class)).isNotNull(); + assertThat(context.getBean(ApiProviderRegistry.class) + .find(cn.lypi.contracts.model.ApiStyle.OPENAI_COMPATIBLE)).isPresent(); + List adapters = context.getBean("openAiCompatibleProviderAdapters", List.class); + OpenAiCompatibleProviderAdapter adapter = adapters.stream() + .map(OpenAiCompatibleProviderAdapter.class::cast) + .filter(candidate -> config(candidate).provider().equals("login-fixture")) + .findFirst() + .orElseThrow(); + assertThat(config(adapter).compat().get("requires-reasoning-content-on-assistant-messages")) + .isIn(true, "true"); + }); + } + @Test void doesNotTriggerRemoteDiscoveryWhenDisabled() { new ApplicationContextRunner() @@ -196,9 +295,210 @@ void configuredModelDescriptorOverridesRemoteAndBuiltInDescriptors() { "lypi.ai.providers.openai.models[1].max-output-tokens=8192" ) .run(context -> { - ModelDescriptor descriptor = openAiModel(context.getBean(ModelRegistry.class), "gpt-5-mini"); + ModelRegistry registry = context.getBean(ModelRegistry.class); + ModelDescriptor descriptor = openAiModel(registry, "gpt-5-mini"); assertThat(descriptor.contextWindow()).isEqualTo(64_000); + assertThat(registry.list()) + .filteredOn(model -> model.provider().equals("openai")) + .extracting(ModelDescriptor::modelId) + .containsExactly("gpt-5-mini"); + }); + } + + @Test + void appliesConfiguredDiscoveryDefaultsOnlyToMissingRemoteMetadata() { + new ApplicationContextRunner() + .withUserConfiguration(LyPiAiAutoConfiguration.class) + .withBean(RemoteModelDiscoveryClient.class, () -> new FixedRemoteModelDiscoveryClient(List.of( + new DiscoveredModel( + "remote-explicit", + OptionalInt.of(128_000), + OptionalInt.of(16_384), + Optional.of(false), + Optional.of(false) + ), + DiscoveredModel.idOnly("remote-defaulted") + ))) + .withPropertyValues( + "lypi.ai.model-discovery.defaults.context-window=192000", + "lypi.ai.model-discovery.defaults.max-output-tokens=12288", + "lypi.ai.model-discovery.defaults.supports-thinking=true", + "lypi.ai.model-discovery.defaults.supports-image-input=false", + "lypi.ai.providers.fixture.enabled=true", + "lypi.ai.providers.fixture.api-style=openai_compatible", + "lypi.ai.providers.fixture.base-url=https://api.fixture.test/v1", + "lypi.ai.providers.fixture.model-discovery.enabled=true" + ) + .run(context -> { + ModelRegistry registry = context.getBean(ModelRegistry.class); + + assertThat(model(registry, "fixture", "remote-explicit")).satisfies(descriptor -> { + assertThat(descriptor.contextWindow()).isEqualTo(128_000); + assertThat(descriptor.maxOutputTokens()).isEqualTo(16_384); + assertThat(descriptor.supportsThinking()).isFalse(); + assertThat(descriptor.supportsImageInput()).isFalse(); + }); + assertThat(model(registry, "fixture", "remote-defaulted")).satisfies(descriptor -> { + assertThat(descriptor.contextWindow()).isEqualTo(192_000); + assertThat(descriptor.maxOutputTokens()).isEqualTo(12_288); + assertThat(descriptor.supportsThinking()).isTrue(); + assertThat(descriptor.supportsImageInput()).isFalse(); + }); + }); + } + + @Test + void passesConfiguredDiscoveryDefaultsToRuntimeProviderLogin() { + Path home = tempDir.resolve("runtime-login-home"); + + new ApplicationContextRunner() + .withUserConfiguration(LyPiAiAutoConfiguration.class) + .withBean(RemoteModelDiscoveryClient.class, () -> new FixedRemoteModelDiscoveryClient("runtime-model")) + .withBean(LoginProviderPropertiesStore.class, () -> new LoginProviderPropertiesStore(home)) + .withPropertyValues( + "lypi.ai.model-discovery.defaults.context-window=192000", + "lypi.ai.model-discovery.defaults.max-output-tokens=12288", + "lypi.ai.model-discovery.defaults.supports-thinking=true", + "lypi.ai.model-discovery.defaults.supports-image-input=false", + "lypi.ai.providers.openai.enabled=false" + ) + .run(context -> { + ProviderLoginResult result = context.getBean(ProviderLoginPort.class).register( + "zen", + "https://api.fixture.test/v1", + "fixture-key" + ); + + assertThat(result.models()).singleElement().satisfies(descriptor -> { + assertThat(descriptor.contextWindow()).isEqualTo(192_000); + assertThat(descriptor.maxOutputTokens()).isEqualTo(12_288); + assertThat(descriptor.supportsThinking()).isTrue(); + assertThat(descriptor.supportsImageInput()).isFalse(); + }); + }); + } + + @Test + void discoveredModelsAreAuthoritativeWhileMatchingLocalMetadataOverridesDefaults() { + new ApplicationContextRunner() + .withUserConfiguration(LyPiAiAutoConfiguration.class) + .withBean(RemoteModelDiscoveryClient.class, () -> new FixedRemoteModelDiscoveryClient(List.of( + new DiscoveredModel( + "remote-a", + OptionalInt.of(128_000), + OptionalInt.of(16_384), + Optional.of(false), + Optional.of(false) + ) + ))) + .withPropertyValues( + "lypi.ai.providers.fixture.enabled=true", + "lypi.ai.providers.fixture.api-style=openai_compatible", + "lypi.ai.providers.fixture.base-url=https://api.fixture.test/v1", + "lypi.ai.providers.fixture.model-discovery.enabled=true", + "lypi.ai.providers.fixture.models[0].model-id=remote-a", + "lypi.ai.providers.fixture.models[0].context-window=96000", + "lypi.ai.providers.fixture.models[0].max-output-tokens=8192", + "lypi.ai.providers.fixture.models[0].supports-thinking=true", + "lypi.ai.providers.fixture.models[0].supports-image-input=true", + "lypi.ai.providers.fixture.models[1].model-id=local-only", + "lypi.ai.providers.fixture.models[1].context-window=64000", + "lypi.ai.providers.fixture.models[1].max-output-tokens=4096" + ) + .run(context -> { + ModelRegistry registry = context.getBean(ModelRegistry.class); + + assertThat(registry.list()) + .filteredOn(model -> model.provider().equals("fixture")) + .extracting(ModelDescriptor::modelId) + .containsExactly("remote-a"); + assertThat(model(registry, "fixture", "remote-a")).satisfies(descriptor -> { + assertThat(descriptor.contextWindow()).isEqualTo(96_000); + assertThat(descriptor.maxOutputTokens()).isEqualTo(8_192); + assertThat(descriptor.supportsThinking()).isTrue(); + assertThat(descriptor.supportsImageInput()).isTrue(); + }); + }); + } + + @Test + void rejectsNonPositiveDiscoveryDefaultTokenLimits() { + new ApplicationContextRunner() + .withUserConfiguration(LyPiAiAutoConfiguration.class) + .withPropertyValues("lypi.ai.model-discovery.defaults.context-window=0") + .run(context -> { + assertThat(context).hasFailed(); + assertThat(rootCause(context.getStartupFailure())) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("Model discovery default token limits must be positive."); + }); + } + + @Test + void discoversEachProviderOnceWhileCreatingTheRegistry() { + CountingRemoteModelDiscoveryClient discovery = new CountingRemoteModelDiscoveryClient("remote-a"); + + new ApplicationContextRunner() + .withUserConfiguration(LyPiAiAutoConfiguration.class) + .withBean(RemoteModelDiscoveryClient.class, () -> discovery) + .withPropertyValues( + "lypi.ai.providers.fixture.enabled=true", + "lypi.ai.providers.fixture.api-style=openai_compatible", + "lypi.ai.providers.fixture.base-url=https://api.fixture.test/v1", + "lypi.ai.providers.fixture.model-discovery.enabled=true" + ) + .run(context -> { + assertThat(context).hasSingleBean(ModelRegistry.class); + assertThat(discovery.calls()).isOne(); + }); + } + + @Test + void failsStartupWhenRemoteModelDiscoveryIsUnavailable() { + new ApplicationContextRunner() + .withUserConfiguration(LyPiAiAutoConfiguration.class) + .withBean(RemoteModelDiscoveryClient.class, FailingRemoteModelDiscoveryClient::new) + .withPropertyValues( + "lypi.ai.providers.fixture.enabled=true", + "lypi.ai.providers.fixture.api-style=openai_compatible", + "lypi.ai.providers.fixture.base-url=https://api.fixture.test/v1", + "lypi.ai.providers.fixture.model-discovery.enabled=true" + ) + .run(context -> { + assertThat(context).hasFailed(); + assertThat(rootCause(context.getStartupFailure())) + .isInstanceOfSatisfying(ModelProviderException.class, error -> + assertThat(error.errorId()).isEqualTo("model.discovery_unavailable")); + }); + } + + @Test + void configuresDiscoveredCompatibleProviderForChatCompletionsSseOnly() { + new ApplicationContextRunner() + .withUserConfiguration(LyPiAiAutoConfiguration.class) + .withBean(RemoteModelDiscoveryClient.class, () -> new FixedRemoteModelDiscoveryClient("remote-a")) + .withPropertyValues( + "lypi.ai.providers.fixture.enabled=true", + "lypi.ai.providers.fixture.api-style=openai_compatible", + "lypi.ai.providers.fixture.base-url=https://api.fixture.test/v1", + "lypi.ai.providers.fixture.api-key=${LYPI_FIXTURE_TOKEN}", + "lypi.ai.providers.fixture.request-style=chat_completions", + "lypi.ai.providers.fixture.fallback-request-style=chat_completions", + "lypi.ai.providers.fixture.transport=sse", + "lypi.ai.providers.fixture.model-discovery.enabled=true" + ) + .run(context -> { + List adapters = context.getBean("openAiCompatibleProviderAdapters", List.class); + OpenAiCompatibleProviderAdapter adapter = adapters.stream() + .map(OpenAiCompatibleProviderAdapter.class::cast) + .filter(candidate -> config(candidate).provider().equals("fixture")) + .findFirst() + .orElseThrow(); + + assertThat(config(adapter).requestStyle()).isEqualTo(RequestStyle.CHAT_COMPLETIONS); + assertThat(config(adapter).fallbackRequestStyle()).isEqualTo(RequestStyle.CHAT_COMPLETIONS); + assertThat(config(adapter).transportMode()).isEqualTo(TransportMode.SSE); }); } @@ -381,21 +681,56 @@ void bindsProviderPropertiesFromYamlResources() { private static final class ThrowingRemoteModelDiscoveryClient extends RemoteModelDiscoveryClient { @Override - public List discover(URI baseUrl, String apiKey, List paths, Duration timeout) { + public List discoverModels(URI baseUrl, String apiKey, List paths, Duration timeout) { throw new AssertionError("Remote discovery should not be called when disabled."); } } private static final class FixedRemoteModelDiscoveryClient extends RemoteModelDiscoveryClient { - private final String modelId; + private final List models; private FixedRemoteModelDiscoveryClient(String modelId) { + this(List.of(DiscoveredModel.idOnly(modelId))); + } + + private FixedRemoteModelDiscoveryClient(List models) { + this.models = List.copyOf(models); + } + + @Override + public List discoverModels(URI baseUrl, String apiKey, List paths, Duration timeout) { + return models; + } + } + + private static final class CountingRemoteModelDiscoveryClient extends RemoteModelDiscoveryClient { + private final AtomicInteger calls = new AtomicInteger(); + private final String modelId; + + private CountingRemoteModelDiscoveryClient(String modelId) { this.modelId = modelId; } @Override - public List discover(URI baseUrl, String apiKey, List paths, Duration timeout) { - return List.of(modelId); + public List discoverModels(URI baseUrl, String apiKey, List paths, Duration timeout) { + calls.incrementAndGet(); + return List.of(DiscoveredModel.idOnly(modelId)); + } + + private int calls() { + return calls.get(); + } + } + + private static final class FailingRemoteModelDiscoveryClient extends RemoteModelDiscoveryClient { + @Override + public List discoverModels(URI baseUrl, String apiKey, List paths, Duration timeout) { + throw new ModelProviderException( + "model.discovery_unavailable", + ErrorSeverity.ERROR, + false, + "Remote model discovery returned no usable models." + ); } } @@ -420,11 +755,23 @@ private static AnthropicProviderConfig anthropicConfig(AnthropicCompatibleProvid } private static ModelDescriptor openAiModel(ModelRegistry registry, String modelId) { + return model(registry, "openai", modelId); + } + + private static ModelDescriptor model(ModelRegistry registry, String provider, String modelId) { return registry.list().stream() - .filter(descriptor -> descriptor.provider().equals("openai")) + .filter(descriptor -> descriptor.provider().equals(provider)) .filter(descriptor -> descriptor.modelId().equals(modelId)) .findFirst() - .orElseThrow(() -> new AssertionError("Missing openai model: " + modelId)); + .orElseThrow(() -> new AssertionError("Missing model: " + provider + "/" + modelId)); + } + + private static Throwable rootCause(Throwable failure) { + Throwable current = failure; + while (current != null && current.getCause() != null) { + current = current.getCause(); + } + return current; } } diff --git a/lypi-boot/src/test/java/cn/lypi/boot/ai/OpenAiCompatibleLoginRealEndToEndTest.java b/lypi-boot/src/test/java/cn/lypi/boot/ai/OpenAiCompatibleLoginRealEndToEndTest.java new file mode 100644 index 00000000..54899410 --- /dev/null +++ b/lypi-boot/src/test/java/cn/lypi/boot/ai/OpenAiCompatibleLoginRealEndToEndTest.java @@ -0,0 +1,460 @@ +package cn.lypi.boot.ai; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.junit.jupiter.api.Assumptions.assumeTrue; + +import cn.lypi.ai.DefaultApiProviderRegistry; +import cn.lypi.ai.DefaultModelPort; +import cn.lypi.ai.DefaultModelRegistry; +import cn.lypi.ai.ProviderAdapterApiProvider; +import cn.lypi.ai.RuntimeModelRegistry; +import cn.lypi.ai.model.DiscoveredModelDefaults; +import cn.lypi.ai.model.RemoteModelDiscoveryClient; +import cn.lypi.contracts.common.JsonSchema; +import cn.lypi.contracts.context.AgentMessage; +import cn.lypi.contracts.context.AttachmentContentBlock; +import cn.lypi.contracts.context.ContentBlock; +import cn.lypi.contracts.context.ContextBudget; +import cn.lypi.contracts.context.ContextSnapshot; +import cn.lypi.contracts.context.MessageKind; +import cn.lypi.contracts.context.MessageRole; +import cn.lypi.contracts.context.TextContentBlock; +import cn.lypi.contracts.context.ThinkingContentBlock; +import cn.lypi.contracts.context.ToolCallContentBlock; +import cn.lypi.contracts.context.ToolResultContentBlock; +import cn.lypi.contracts.model.ApiStyle; +import cn.lypi.contracts.model.AssistantDone; +import cn.lypi.contracts.model.AssistantEventStream; +import cn.lypi.contracts.model.AssistantStreamEvent; +import cn.lypi.contracts.model.CostProfile; +import cn.lypi.contracts.model.ModelDescriptor; +import cn.lypi.contracts.model.ModelSelection; +import cn.lypi.contracts.model.TextDelta; +import cn.lypi.contracts.model.ThinkingDelta; +import cn.lypi.contracts.model.ThinkingLevel; +import cn.lypi.contracts.model.ToolCallDelta; +import cn.lypi.contracts.prompt.SystemPrompt; +import cn.lypi.contracts.runtime.ProviderLoginResult; +import cn.lypi.contracts.security.AgentMode; +import cn.lypi.contracts.security.PermissionMode; +import cn.lypi.contracts.tool.ToolDescriptor; +import cn.lypi.contracts.tool.ToolRegistrySnapshot; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import java.io.IOException; +import java.io.InputStream; +import java.math.BigDecimal; +import java.net.URI; +import java.nio.file.Files; +import java.nio.file.Path; +import java.time.Instant; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.Properties; +import java.util.concurrent.TimeUnit; +import java.util.stream.Collectors; +import java.util.stream.StreamSupport; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; +import org.junit.jupiter.api.io.TempDir; + +class OpenAiCompatibleLoginRealEndToEndTest { + private static final URI BASE_URL = URI.create(System.getProperty( + "lypi.openai-compatible.e2e.base-url", + "https://opencode.ai/zen/go/v1" + )); + private static final String CHANNEL = "compat-live"; + private static final String THINKING_MODEL = System.getProperty( + "lypi.openai-compatible.e2e.thinking-model", + "deepseek-v4-flash" + ); + private static final String IMAGE_MODEL = System.getProperty( + "lypi.openai-compatible.e2e.image-model", + "kimi-k2.6" + ); + private static final String TOOL_NAME = "lypi_compatibility_probe"; + private static final String TOOL_TOKEN = "LIVE_TOOL_OK"; + private static final String IMAGE_TOKEN = "RED_BLUE"; + private static final String RED_BLUE_PNG = "data:image/png;base64," + + "iVBORw0KGgoAAAANSUhEUgAAAAIAAAABCAIAAAB7QOjdAAAAD0lEQVR4nGP4z8DAwPAfAAcAAf9+CLHQAAAAAElFTkSuQmCC"; + + @TempDir + Path tempDir; + + @Test + @Timeout(value = 300, unit = TimeUnit.SECONDS) + void logsInAndStreamsThinkingToolContinuationAndImageChatCompletions() { + assumeTrue( + Boolean.getBoolean("lypi.openai-compatible.e2e"), + "Enable with -Dlypi.openai-compatible.e2e=true" + ); + String apiKey = apiKey(); + RuntimeModelRegistry modelRegistry = new DefaultModelRegistry(List.of()); + ProviderAdapterApiProvider dispatcher = new ProviderAdapterApiProvider( + ApiStyle.OPENAI_COMPATIBLE, + List.of() + ); + Path home = tempDir.resolve("home"); + OpenAiCompatibleProviderLoginService loginService = new OpenAiCompatibleProviderLoginService( + new RemoteModelDiscoveryClient(), + modelRegistry, + dispatcher, + new LoginProviderPropertiesStore(home), + discoveredDefaults() + ); + + ProviderLoginResult login = safely( + "OpenAI-compatible login", + () -> loginService.register(CHANNEL, BASE_URL.toString(), apiKey) + ); + assertThat(login.provider()).isEqualTo(CHANNEL); + assertThat(login.models()).isNotEmpty(); + assertThat(login.models()).allSatisfy(model -> { + assertThat(model.provider()).isEqualTo(CHANNEL); + assertThat(model.contextWindow()).isPositive(); + assertThat(model.maxOutputTokens()).isPositive(); + }); + assertManagedProperties(home.resolve(".ly-pi/login-providers.properties")); + assertThat(modelRegistry.list()).containsExactlyElementsOf(login.models()); + + ModelDescriptor thinkingDescriptor = requireModel(login, THINKING_MODEL, "thinking and tool"); + assertThat(thinkingDescriptor.supportsThinking()).isTrue(); + ModelDescriptor imageDescriptor = requireModel(login, IMAGE_MODEL, "image"); + assertThat(imageDescriptor.supportsImageInput()).isTrue(); + DefaultModelPort modelPort = new DefaultModelPort( + modelRegistry, + new DefaultApiProviderRegistry(List.of(dispatcher)) + ); + List allEvents = new ArrayList<>(); + + List thinkingEvents = stream( + modelPort, + context( + thinkingDescriptor, + ThinkingLevel.HIGH, + List.of(userMessage( + "thinking-user", + "Think briefly, then answer in one short sentence that contains LIVE_THINKING_OK." + )) + ), + new ToolRegistrySnapshot(List.of()), + "HIGH thinking chat" + ); + allEvents.addAll(thinkingEvents); + assertThat(joinedThinking(thinkingEvents)).isNotBlank(); + assertThat(joinedText(thinkingEvents)).isNotBlank(); + assertCompleted(thinkingEvents); + + ToolRegistrySnapshot tools = toolRegistry(); + AgentMessage toolPrompt = userMessage( + "tool-user", + "Call lypi_compatibility_probe exactly once with value LIVE_TOOL_OK. " + + "Do not answer before calling it. After receiving the tool result, return LIVE_TOOL_OK." + ); + List toolEvents = stream( + modelPort, + context(thinkingDescriptor, ThinkingLevel.HIGH, List.of(toolPrompt)), + tools, + "tool-call chat" + ); + allEvents.addAll(toolEvents); + String toolThinking = joinedThinking(toolEvents); + assertThat(toolThinking).isNotBlank(); + ToolCallDelta toolCall = completeToolCall(toolEvents); + assertThat(toolCall.toolName()).isEqualTo(TOOL_NAME); + assertThat(toolCall.toolUseId()).isNotBlank(); + assertThat(toolCall.partialInput()).containsEntry("value", TOOL_TOKEN); + assertCompleted(toolEvents); + + AgentMessage assistantToolCall = assistantToolCall(toolCall, toolThinking, joinedText(toolEvents)); + AgentMessage toolResult = toolResult(toolCall.toolUseId()); + List continuationEvents = stream( + modelPort, + context( + thinkingDescriptor, + ThinkingLevel.HIGH, + List.of(toolPrompt, assistantToolCall, toolResult) + ), + new ToolRegistrySnapshot(List.of()), + "tool-result continuation chat" + ); + allEvents.addAll(continuationEvents); + assertThat(joinedText(continuationEvents)).contains(TOOL_TOKEN); + assertCompleted(continuationEvents); + + ContextSnapshot imageContext = context( + imageDescriptor, + ThinkingLevel.OFF, + List.of(imageMessage()) + ); + List imageEvents = stream( + modelPort, + imageContext, + new ToolRegistrySnapshot(List.of()), + "image chat" + ); + allEvents.addAll(imageEvents); + assertThat(joinedText(imageEvents)).contains(IMAGE_TOKEN); + assertCompleted(imageEvents); + + String eventTypes = allEvents.stream() + .map(event -> event.getClass().getSimpleName()) + .distinct() + .sorted() + .collect(Collectors.joining(",")); + System.out.printf( + "OpenAI-compatible E2E models=%d thinkingModel=%s toolModel=%s imageModel=%s events=%s%n", + login.models().size(), + THINKING_MODEL, + THINKING_MODEL, + IMAGE_MODEL, + eventTypes + ); + } + + private static String apiKey() { + String configured = System.getProperty("lypi.openai-compatible.e2e.api-key"); + if (configured != null && !configured.isBlank()) { + return configured; + } + String authEntry = System.getProperty( + "lypi.openai-compatible.e2e.auth-entry", + "opencode-go" + ); + Path authFile = Path.of(System.getProperty("user.home"), ".pi", "agent", "auth.json"); + try { + JsonNode credential = new ObjectMapper().readTree(authFile.toFile()).path(authEntry); + String apiKey = credential.path("key").asText(); + if (!apiKey.isBlank()) { + return apiKey; + } + } catch (IOException | RuntimeException error) { + throw new AssertionError("Unable to load the configured E2E credential."); + } + throw new AssertionError("The configured E2E credential is missing an API key."); + } + + private static DiscoveredModelDefaults discoveredDefaults() { + return new DiscoveredModelDefaults( + DiscoveredModelDefaults.DEFAULT_CONTEXT_WINDOW, + DiscoveredModelDefaults.DEFAULT_MAX_OUTPUT_TOKENS, + DiscoveredModelDefaults.DEFAULT_SUPPORTS_THINKING, + DiscoveredModelDefaults.DEFAULT_SUPPORTS_IMAGE_INPUT, + new CostProfile(BigDecimal.ZERO, BigDecimal.ZERO, "USD"), + Map.of() + ); + } + + private static ModelDescriptor requireModel( + ProviderLoginResult login, + String modelId, + String purpose + ) { + return login.models().stream() + .filter(model -> model.modelId().equals(modelId)) + .findFirst() + .orElseThrow(() -> new AssertionError( + "Required " + purpose + " E2E model is unavailable: " + modelId + )); + } + + private static void assertManagedProperties(Path storeFile) { + assertThat(storeFile).exists(); + Properties properties = new Properties(); + try (InputStream input = Files.newInputStream(storeFile)) { + properties.load(input); + } catch (IOException error) { + throw new AssertionError("Unable to inspect managed E2E provider properties."); + } + boolean cachesModels = properties.stringPropertyNames().stream() + .anyMatch(name -> name.contains(".models[")); + boolean disablesThinking = properties.stringPropertyNames().stream() + .filter(name -> name.endsWith(".supports-thinking")) + .map(properties::getProperty) + .anyMatch("false"::equalsIgnoreCase); + assertThat(cachesModels).as("managed login properties cache no model list").isFalse(); + assertThat(disablesThinking).as("managed login properties do not disable thinking").isFalse(); + } + + private static ContextSnapshot context( + ModelDescriptor descriptor, + ThinkingLevel thinkingLevel, + List messages + ) { + int autoCompactThreshold = (int) Math.min( + Integer.MAX_VALUE, + Math.max(1L, descriptor.contextWindow() * 4L / 5L) + ); + return new ContextSnapshot( + new SystemPrompt( + "Follow the user's compatibility-test instructions exactly.", + List.of("e2e"), + "e2e" + ), + messages, + new ModelSelection(CHANNEL, descriptor.modelId(), thinkingLevel), + thinkingLevel, + AgentMode.EXECUTE, + PermissionMode.ASK, + new ContextBudget( + 0, + descriptor.contextWindow(), + autoCompactThreshold, + descriptor.maxOutputTokens(), + Math.min(8_192, descriptor.maxOutputTokens()), + 0, + 0, + BigDecimal.ZERO + ) + ); + } + + private static AgentMessage userMessage(String id, String text) { + return message(id, MessageRole.USER, MessageKind.TEXT, List.of(new TextContentBlock(text))); + } + + private static AgentMessage imageMessage() { + return message( + "image-user", + MessageRole.USER, + MessageKind.ATTACHMENT, + List.of( + new TextContentBlock( + "Inspect the attached two-pixel image. If its left pixel is red and its right pixel is blue, " + + "reply with exactly RED_BLUE." + ), + new AttachmentContentBlock( + "red-blue-image", + "two-pixel red and blue image", + "image/png", + Map.of("imageUrl", RED_BLUE_PNG, "detail", "high") + ) + ) + ); + } + + private static AgentMessage assistantToolCall( + ToolCallDelta toolCall, + String thinking, + String text + ) { + List content = new ArrayList<>(); + if (!thinking.isBlank()) { + content.add(new ThinkingContentBlock(thinking)); + } + if (!text.isBlank()) { + content.add(new TextContentBlock(text)); + } + content.add(new ToolCallContentBlock( + toolCall.toolUseId(), + toolCall.toolName(), + "", + Map.of("input", toolCall.partialInput()) + )); + return message("tool-assistant", MessageRole.ASSISTANT, MessageKind.TOOL_CALL, content); + } + + private static AgentMessage toolResult(String toolUseId) { + return message( + "tool-result", + MessageRole.TOOL_RESULT, + MessageKind.TOOL_RESULT, + List.of(new ToolResultContentBlock(toolUseId, TOOL_TOKEN, false)) + ); + } + + private static AgentMessage message( + String id, + MessageRole role, + MessageKind kind, + List content + ) { + return new AgentMessage( + id, + role, + kind, + content, + Instant.EPOCH, + Optional.empty(), + Optional.empty() + ); + } + + private static ToolRegistrySnapshot toolRegistry() { + return new ToolRegistrySnapshot(List.of(new ToolDescriptor( + TOOL_NAME, + List.of(), + "Required compatibility probe. Call it with the exact requested value.", + new JsonSchema(Map.of( + "type", "object", + "properties", Map.of("value", Map.of( + "type", "string", + "enum", List.of(TOOL_TOKEN) + )), + "required", List.of("value"), + "additionalProperties", false + )), + true, + false + ))); + } + + private static List stream( + DefaultModelPort modelPort, + ContextSnapshot context, + ToolRegistrySnapshot tools, + String operation + ) { + return safely(operation, () -> { + try (AssistantEventStream stream = modelPort.stream(context, tools, () -> false)) { + return StreamSupport.stream(stream.spliterator(), false).toList(); + } + }); + } + + private static ToolCallDelta completeToolCall(List events) { + return events.stream() + .filter(ToolCallDelta.class::isInstance) + .map(ToolCallDelta.class::cast) + .filter(ToolCallDelta::complete) + .reduce((first, second) -> second) + .orElseThrow(() -> new AssertionError("Tool-call chat emitted no complete tool call.")); + } + + private static String joinedThinking(List events) { + return events.stream() + .filter(ThinkingDelta.class::isInstance) + .map(ThinkingDelta.class::cast) + .map(ThinkingDelta::text) + .collect(Collectors.joining()); + } + + private static String joinedText(List events) { + return events.stream() + .filter(TextDelta.class::isInstance) + .map(TextDelta.class::cast) + .map(TextDelta::text) + .collect(Collectors.joining()); + } + + private static void assertCompleted(List events) { + assertThat(events.stream().anyMatch(AssistantDone.class::isInstance)) + .as("provider stream emitted AssistantDone") + .isTrue(); + } + + private static T safely(String operation, UnsafeSupplier supplier) { + try { + return supplier.get(); + } catch (RuntimeException error) { + throw new AssertionError(operation + " failed (" + error.getClass().getSimpleName() + ")."); + } + } + + @FunctionalInterface + private interface UnsafeSupplier { + T get(); + } +} diff --git a/lypi-boot/src/test/java/cn/lypi/boot/ai/OpenAiCompatibleProviderLoginServiceTest.java b/lypi-boot/src/test/java/cn/lypi/boot/ai/OpenAiCompatibleProviderLoginServiceTest.java new file mode 100644 index 00000000..b2044770 --- /dev/null +++ b/lypi-boot/src/test/java/cn/lypi/boot/ai/OpenAiCompatibleProviderLoginServiceTest.java @@ -0,0 +1,558 @@ +package cn.lypi.boot.ai; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import cn.lypi.ai.DefaultModelRegistry; +import cn.lypi.ai.ProviderAdapterApiProvider; +import cn.lypi.ai.RuntimeModelRegistry; +import cn.lypi.ai.model.DiscoveredModel; +import cn.lypi.ai.model.DiscoveredModelDefaults; +import cn.lypi.ai.model.RemoteModelDiscoveryClient; +import cn.lypi.ai.provider.ProviderRequest; +import cn.lypi.ai.provider.RequestStyle; +import cn.lypi.ai.provider.TransportMode; +import cn.lypi.ai.provider.openai.OpenAiCompatibleProviderAdapter; +import cn.lypi.ai.provider.openai.OpenAiProviderConfig; +import cn.lypi.contracts.error.ModelProviderException; +import cn.lypi.contracts.model.ApiStyle; +import cn.lypi.contracts.model.CostProfile; +import cn.lypi.contracts.model.ModelDescriptor; +import cn.lypi.contracts.runtime.ProviderLoginResult; +import com.sun.net.httpserver.HttpExchange; +import com.sun.net.httpserver.HttpServer; +import java.io.IOException; +import java.lang.reflect.Field; +import java.net.InetSocketAddress; +import java.net.URI; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.attribute.PosixFileAttributeView; +import java.nio.file.attribute.PosixFilePermission; +import java.util.List; +import java.util.Map; +import java.util.Properties; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +class OpenAiCompatibleProviderLoginServiceTest { + private static final String AUTH_KEY = "test-auth-key"; + + @TempDir + Path tempDir; + + private HttpServer server; + + @AfterEach + void stopServer() { + if (server != null) { + server.stop(0); + } + } + + @Test + void persistsAndRegistersOnlyAfterModelDiscoverySucceeds() throws Exception { + AtomicReference authorization = new AtomicReference<>(); + startServer(exchange -> { + authorization.set(exchange.getRequestHeaders().getFirst("Authorization")); + respond(exchange, 200, """ + {"data":[ + { + "id":"zeta", + "context_window":128000, + "max_output_tokens":16384, + "supports_reasoning":false, + "supports_image_input":false + }, + {"id":"alpha"} + ]} + """); + }); + RuntimeModelRegistry registry = new DefaultModelRegistry(List.of()); + ProviderAdapterApiProvider dispatcher = new ProviderAdapterApiProvider(ApiStyle.OPENAI_COMPATIBLE, List.of()); + Path home = tempDir.resolve("home"); + OpenAiCompatibleProviderLoginService service = service(registry, dispatcher, home); + + ProviderLoginResult result = service.register("zen", baseUrl().toString() + "/", AUTH_KEY); + Path storeFile = home.resolve(".ly-pi/login-providers.properties"); + + assertThat(("Bearer " + AUTH_KEY).equals(authorization.get())).isTrue(); + assertThat(result.provider()).isEqualTo("zen"); + assertThat(result.models()).extracting(ModelDescriptor::modelId).containsExactly("zeta", "alpha"); + assertThat(result.models().get(0)).satisfies(model -> { + assertThat(model.provider()).isEqualTo("zen"); + assertThat(model.baseUrl()).isEqualTo(baseUrl()); + assertThat(model.contextWindow()).isEqualTo(128_000); + assertThat(model.maxOutputTokens()).isEqualTo(16_384); + assertThat(model.supportsThinking()).isFalse(); + assertThat(model.supportsImageInput()).isFalse(); + }); + assertThat(result.models().get(1)).satisfies(model -> { + assertThat(model.provider()).isEqualTo("zen"); + assertThat(model.contextWindow()).isEqualTo(256_000); + assertThat(model.maxOutputTokens()).isEqualTo(8_192); + assertThat(model.supportsThinking()).isTrue(); + assertThat(model.supportsImageInput()).isTrue(); + }); + assertThat(registry.list()).containsExactlyElementsOf(result.models()); + assertThat(storeFile).exists(); + Properties stored = properties(storeFile); + assertThat(stored.stringPropertyNames()) + .contains( + "lypi.ai.providers." + result.provider() + ".request-style", + "lypi.ai.providers." + result.provider() + ".fallback-request-style", + "lypi.ai.providers." + result.provider() + ".transport" + ); + assertThat(stored.getProperty( + "lypi.ai.providers." + result.provider() + ".request-style" + )).isEqualTo("chat_completions"); + assertThat(stored.getProperty( + "lypi.ai.providers." + result.provider() + ".fallback-request-style" + )).isEqualTo("chat_completions"); + assertThat(stored.getProperty( + "lypi.ai.providers." + result.provider() + ".transport" + )).isEqualTo("sse"); + assertThat(stored.getProperty( + "lypi.ai.providers." + result.provider() + + ".compat.requires-reasoning-content-on-assistant-messages" + )).isEqualTo("true"); + assertThat(stored.stringPropertyNames()) + .noneMatch(name -> name.contains(".models[") || name.endsWith("supports-thinking")); + assertThat(config(dispatcher, result.provider()).requestStyle()).isEqualTo(RequestStyle.CHAT_COMPLETIONS); + assertThat(config(dispatcher, result.provider()).fallbackRequestStyle()).isEqualTo(RequestStyle.CHAT_COMPLETIONS); + assertThat(config(dispatcher, result.provider()).transportMode()).isEqualTo(TransportMode.SSE); + assertThat(config(dispatcher, result.provider()).compat()) + .containsEntry("requires-reasoning-content-on-assistant-messages", true); + assertThat(config(dispatcher, result.provider()).toString().contains(AUTH_KEY)).isFalse(); + assertPrivateFile(storeFile); + } + + @Test + void fallsBackToModelEndpointAndReplacesOnlyTheNamedProvider() throws Exception { + AtomicReference calls = new AtomicReference<>(0); + List requestedPaths = new CopyOnWriteArrayList<>(); + startServer(exchange -> { + requestedPaths.add(exchange.getRequestURI().getPath()); + if (exchange.getRequestURI().getPath().endsWith("/models")) { + respond(exchange, 404, ""); + return; + } + int call = calls.updateAndGet(value -> value + 1); + String modelId = switch (call) { + case 1 -> "old-model"; + case 2 -> "other-model"; + default -> "new-model"; + }; + respond(exchange, 200, "{\"models\":[{\"id\":\"" + modelId + "\"}]}"); + }); + RuntimeModelRegistry registry = new DefaultModelRegistry(List.of()); + ProviderAdapterApiProvider dispatcher = new ProviderAdapterApiProvider(ApiStyle.OPENAI_COMPATIBLE, List.of()); + Path home = tempDir.resolve("home"); + OpenAiCompatibleProviderLoginService service = service(registry, dispatcher, home); + + ProviderLoginResult first = service.register("zen-a", baseUrl().toString(), AUTH_KEY); + OpenAiCompatibleProviderAdapter firstAdapter = adapter(dispatcher, first.provider()); + ProviderLoginResult other = service.register("zen-b", baseUrl().toString(), AUTH_KEY); + OpenAiCompatibleProviderAdapter otherAdapter = adapter(dispatcher, other.provider()); + String replacementKey = "replacement-auth-key"; + ProviderLoginResult second = service.register("zen-a", baseUrl().toString(), replacementKey); + + assertThat(second.provider()).isEqualTo(first.provider()); + assertThat(registry.list()) + .filteredOn(model -> model.provider().equals(first.provider())) + .extracting(ModelDescriptor::modelId) + .containsExactly("new-model"); + assertThat(registry.list()) + .filteredOn(model -> model.provider().equals(other.provider())) + .extracting(ModelDescriptor::modelId) + .containsExactly("other-model"); + Properties stored = properties(home.resolve(".ly-pi/login-providers.properties")); + assertThat(stored.stringPropertyNames()) + .anyMatch(name -> name.startsWith("lypi.ai.providers.zen-a.")) + .anyMatch(name -> name.startsWith("lypi.ai.providers.zen-b.")); + assertThat(stored.getProperty("lypi.ai.providers.zen-a.api-key")).isEqualTo(replacementKey); + assertThat(stored.getProperty("lypi.ai.providers.zen-b.api-key")).isEqualTo(AUTH_KEY); + assertThat(calls.get()).isEqualTo(3); + assertThat(requestedPaths).containsExactly( + "/v1/models", "/v1/model", + "/v1/models", "/v1/model", + "/v1/models", "/v1/model" + ); + assertThat(adapter(dispatcher, second.provider())).isNotSameAs(firstAdapter); + assertThat(adapter(dispatcher, other.provider())).isSameAs(otherAdapter); + assertThat(config(dispatcher, second.provider()).apiKey()).isEqualTo(replacementKey); + assertThat(config(dispatcher, other.provider()).apiKey()).isEqualTo(AUTH_KEY); + } + + @Test + void leavesFileAndRuntimeUnchangedWhenDiscoveryFails() throws Exception { + startServer(exchange -> respond(exchange, 401, "")); + RuntimeModelRegistry registry = new DefaultModelRegistry(List.of(existingModel())); + ProviderAdapterApiProvider dispatcher = new ProviderAdapterApiProvider(ApiStyle.OPENAI_COMPATIBLE, List.of()); + Path home = tempDir.resolve("home"); + OpenAiCompatibleProviderLoginService service = service(registry, dispatcher, home); + + assertThatThrownBy(() -> service.register("zen", baseUrl().toString(), AUTH_KEY)) + .isInstanceOf(ModelProviderException.class) + .hasMessageNotContaining(AUTH_KEY); + + assertThat(registry.list()).containsExactly(existingModel()); + assertThat(home.resolve(".ly-pi/login-providers.properties")).doesNotExist(); + assertThat(adapterCount(dispatcher)).isZero(); + } + + @Test + void leavesFileAndRuntimeUnchangedWhenBothDiscoveryEndpointsHaveNoModels() throws Exception { + List requestedPaths = new CopyOnWriteArrayList<>(); + startServer(exchange -> { + requestedPaths.add(exchange.getRequestURI().getPath()); + respond(exchange, 200, "{\"data\":[]}"); + }); + RuntimeModelRegistry registry = new DefaultModelRegistry(List.of(existingModel())); + ProviderAdapterApiProvider dispatcher = new ProviderAdapterApiProvider(ApiStyle.OPENAI_COMPATIBLE, List.of()); + Path home = tempDir.resolve("home"); + OpenAiCompatibleProviderLoginService service = service(registry, dispatcher, home); + + assertThatThrownBy(() -> service.register("zen", baseUrl().toString(), AUTH_KEY)) + .isInstanceOf(ModelProviderException.class) + .hasMessageNotContaining(AUTH_KEY); + + assertThat(requestedPaths).containsExactly("/v1/models", "/v1/model"); + assertThat(registry.list()).containsExactly(existingModel()); + assertThat(adapterCount(dispatcher)).isZero(); + assertThat(home.resolve(".ly-pi/login-providers.properties")).doesNotExist(); + } + + @Test + void leavesExistingStateUntouchedWhenPersistenceFails() throws Exception { + RuntimeModelRegistry registry = new DefaultModelRegistry(List.of()); + ProviderAdapterApiProvider dispatcher = new ProviderAdapterApiProvider(ApiStyle.OPENAI_COMPATIBLE, List.of()); + Path home = tempDir.resolve("home"); + Path storeFile = home.resolve(".ly-pi/login-providers.properties"); + OpenAiCompatibleProviderLoginService initialService = service( + registry, + dispatcher, + fixedDiscovery("old-model"), + new LoginProviderPropertiesStore(home) + ); + ProviderLoginResult initial = initialService.register( + "zen", + "https://old.example.test/v1", + "old-key" + ); + String initialFile = Files.readString(storeFile); + OpenAiCompatibleProviderAdapter initialAdapter = adapter(dispatcher, "zen"); + + LoginProviderPropertiesStore failingStore = new LoginProviderPropertiesStore(home) { + @Override + public void save(String provider, URI baseUrl, String authKey) throws IOException { + throw new IOException("simulated persistence failure"); + } + }; + OpenAiCompatibleProviderLoginService service = service( + registry, + dispatcher, + fixedDiscovery("verified-model"), + failingStore + ); + + String sensitiveUrl = "https://example.test/private-path"; + + assertThatThrownBy(() -> service.register("zen", sensitiveUrl, AUTH_KEY)) + .isInstanceOf(ModelProviderException.class) + .hasMessageNotContaining("zen") + .hasMessageNotContaining(sensitiveUrl) + .hasMessageNotContaining(AUTH_KEY); + + assertThat(Files.readString(storeFile)).isEqualTo(initialFile); + assertThat(registry.list()).containsExactlyElementsOf(initial.models()); + assertThat(adapter(dispatcher, "zen")).isSameAs(initialAdapter); + assertThat(adapterCount(dispatcher)).isOne(); + } + + @Test + void rejectsInvalidChannelNamesBeforeDiscoveryPersistenceOrRuntimeMutation() { + AtomicInteger discoveryCalls = new AtomicInteger(); + RemoteModelDiscoveryClient discovery = new RemoteModelDiscoveryClient() { + @Override + public List discoverModels( + URI baseUrl, + String apiKey, + List paths, + java.time.Duration timeout + ) { + discoveryCalls.incrementAndGet(); + return List.of(DiscoveredModel.idOnly("unexpected")); + } + }; + RuntimeModelRegistry registry = new DefaultModelRegistry(List.of(existingModel())); + ProviderAdapterApiProvider dispatcher = new ProviderAdapterApiProvider(ApiStyle.OPENAI_COMPATIBLE, List.of()); + Path home = tempDir.resolve("home"); + OpenAiCompatibleProviderLoginService service = service( + registry, + dispatcher, + discovery, + new LoginProviderPropertiesStore(home) + ); + + for (String invalid : List.of("Zen", "with space", ".nested", "a/b", "a".repeat(65))) { + assertThatThrownBy(() -> service.register(invalid, "not-a-url", " ")) + .isInstanceOfSatisfying(ModelProviderException.class, error -> + assertThat(error.errorId()).isEqualTo("provider.login_invalid_channel_name")) + .hasMessage("Provider channel name must match [a-z0-9][a-z0-9_-]{0,63}.") + .hasMessageNotContaining(invalid); + } + assertThatThrownBy(() -> service.register(" ", "not-a-url", " ")) + .isInstanceOfSatisfying(ModelProviderException.class, error -> + assertThat(error.errorId()).isEqualTo("provider.login_invalid_channel_name")) + .hasMessage("Provider channel name must match [a-z0-9][a-z0-9_-]{0,63}."); + assertThatThrownBy(() -> service.register(null, "not-a-url", " ")) + .isInstanceOfSatisfying(ModelProviderException.class, error -> + assertThat(error.errorId()).isEqualTo("provider.login_invalid_channel_name")); + + assertThat(discoveryCalls).hasValue(0); + assertThat(registry.list()).containsExactly(existingModel()); + assertThat(adapterCount(dispatcher)).isZero(); + assertThat(home.resolve(".ly-pi/login-providers.properties")).doesNotExist(); + } + + @Test + void appliesInjectedDefaultsToIdOnlyDiscoveredModels() { + RuntimeModelRegistry registry = new DefaultModelRegistry(List.of()); + ProviderAdapterApiProvider dispatcher = new ProviderAdapterApiProvider(ApiStyle.OPENAI_COMPATIBLE, List.of()); + RemoteModelDiscoveryClient discovery = new RemoteModelDiscoveryClient() { + @Override + public List discoverModels( + URI baseUrl, + String apiKey, + List paths, + java.time.Duration timeout + ) { + return List.of(DiscoveredModel.idOnly("defaulted")); + } + }; + DiscoveredModelDefaults defaults = new DiscoveredModelDefaults( + 192_000, + 12_288, + true, + false, + new CostProfile(java.math.BigDecimal.ZERO, java.math.BigDecimal.ZERO, "USD"), + Map.of() + ); + OpenAiCompatibleProviderLoginService service = new OpenAiCompatibleProviderLoginService( + discovery, + registry, + dispatcher, + new LoginProviderPropertiesStore(tempDir.resolve("home")), + defaults + ); + + ModelDescriptor descriptor = service.register( + "zen", + "https://example.test/v1", + AUTH_KEY + ).models().getFirst(); + + assertThat(descriptor.contextWindow()).isEqualTo(192_000); + assertThat(descriptor.maxOutputTokens()).isEqualTo(12_288); + assertThat(descriptor.supportsThinking()).isTrue(); + assertThat(descriptor.supportsImageInput()).isFalse(); + } + + @Test + void rejectsInvalidUrlAndBlankKeyBeforeMutatingRuntime() { + RuntimeModelRegistry registry = new DefaultModelRegistry(List.of(existingModel())); + ProviderAdapterApiProvider dispatcher = new ProviderAdapterApiProvider(ApiStyle.OPENAI_COMPATIBLE, List.of()); + OpenAiCompatibleProviderLoginService service = service(registry, dispatcher, tempDir.resolve("home")); + + for (String invalidUrl : List.of( + "https://user:pass@example.test/v1", + "https://example.test/v1?tenant=test", + "https://example.test/v1#fragment", + "ftp://example.test/v1" + )) { + assertThatThrownBy(() -> service.register("zen", invalidUrl, AUTH_KEY)) + .isInstanceOf(ModelProviderException.class) + .hasMessageNotContaining(invalidUrl); + } + assertThatThrownBy(() -> service.register("zen", "https://example.test/v1", " ")) + .isInstanceOf(ModelProviderException.class) + .hasMessageNotContaining(AUTH_KEY); + + assertThat(registry.list()).containsExactly(existingModel()); + assertThat(adapterCount(dispatcher)).isZero(); + } + + @Test + void redactsAuthKeyFromDiscoveryFailuresAndProviderRequestStrings() { + RuntimeModelRegistry registry = new DefaultModelRegistry(List.of(existingModel())); + ProviderAdapterApiProvider dispatcher = new ProviderAdapterApiProvider(ApiStyle.OPENAI_COMPATIBLE, List.of()); + RemoteModelDiscoveryClient unsafeDiscovery = new RemoteModelDiscoveryClient() { + @Override + public List discoverModels( + URI baseUrl, + String apiKey, + List paths, + java.time.Duration timeout + ) { + throw new ModelProviderException( + "test.unsafe_discovery", + cn.lypi.contracts.error.ErrorSeverity.ERROR, + false, + "unsafe discovery " + apiKey + ); + } + }; + OpenAiCompatibleProviderLoginService service = service( + registry, + dispatcher, + unsafeDiscovery, + new LoginProviderPropertiesStore(tempDir.resolve("home")) + ); + ProviderRequest request = new ProviderRequest( + URI.create("https://example.test/v1/chat/completions"), + Map.of("Authorization", "Bearer " + AUTH_KEY), + "{}" + ); + + assertThatThrownBy(() -> service.register("zen", "https://example.test/v1", AUTH_KEY)) + .isInstanceOf(ModelProviderException.class) + .hasMessageNotContaining(AUTH_KEY); + assertThat(request.toString().contains(AUTH_KEY)).isFalse(); + assertThat(registry.list()).containsExactly(existingModel()); + assertThat(adapterCount(dispatcher)).isZero(); + } + + private OpenAiCompatibleProviderLoginService service( + RuntimeModelRegistry registry, + ProviderAdapterApiProvider dispatcher, + Path home + ) { + return service(registry, dispatcher, new RemoteModelDiscoveryClient(), new LoginProviderPropertiesStore(home)); + } + + private OpenAiCompatibleProviderLoginService service( + RuntimeModelRegistry registry, + ProviderAdapterApiProvider dispatcher, + RemoteModelDiscoveryClient discovery, + LoginProviderPropertiesStore store + ) { + return new OpenAiCompatibleProviderLoginService(discovery, registry, dispatcher, store, defaults()); + } + + private static DiscoveredModelDefaults defaults() { + return new DiscoveredModelDefaults( + DiscoveredModelDefaults.DEFAULT_CONTEXT_WINDOW, + DiscoveredModelDefaults.DEFAULT_MAX_OUTPUT_TOKENS, + DiscoveredModelDefaults.DEFAULT_SUPPORTS_THINKING, + DiscoveredModelDefaults.DEFAULT_SUPPORTS_IMAGE_INPUT, + new CostProfile(java.math.BigDecimal.ZERO, java.math.BigDecimal.ZERO, "USD"), + Map.of() + ); + } + + private static RemoteModelDiscoveryClient fixedDiscovery(String modelId) { + return new RemoteModelDiscoveryClient() { + @Override + public List discoverModels( + URI baseUrl, + String apiKey, + List paths, + java.time.Duration timeout + ) { + return List.of(DiscoveredModel.idOnly(modelId)); + } + }; + } + + private URI baseUrl() { + return URI.create("http://127.0.0.1:" + server.getAddress().getPort() + "/v1"); + } + + private void startServer(ExchangeHandler handler) throws IOException { + server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); + server.createContext("/v1/models", handler::handle); + server.createContext("/v1/model", handler::handle); + server.start(); + } + + private static Properties properties(Path path) throws IOException { + Properties properties = new Properties(); + try (java.io.InputStream input = Files.newInputStream(path)) { + properties.load(input); + } + return properties; + } + + private static OpenAiProviderConfig config(ProviderAdapterApiProvider dispatcher, String provider) { + try { + OpenAiCompatibleProviderAdapter adapter = adapter(dispatcher, provider); + Field configField = OpenAiCompatibleProviderAdapter.class.getDeclaredField("config"); + configField.setAccessible(true); + return (OpenAiProviderConfig) configField.get(adapter); + } catch (ReflectiveOperationException error) { + throw new AssertionError("Unable to inspect registered provider adapter", error); + } + } + + private static OpenAiCompatibleProviderAdapter adapter(ProviderAdapterApiProvider dispatcher, String provider) { + return (OpenAiCompatibleProviderAdapter) adapters(dispatcher).get(provider); + } + + private static int adapterCount(ProviderAdapterApiProvider dispatcher) { + return adapters(dispatcher).size(); + } + + @SuppressWarnings("unchecked") + private static Map adapters(ProviderAdapterApiProvider dispatcher) { + try { + Field adaptersField = ProviderAdapterApiProvider.class.getDeclaredField("adapters"); + adaptersField.setAccessible(true); + java.util.concurrent.atomic.AtomicReference> adapters = + (java.util.concurrent.atomic.AtomicReference>) adaptersField.get(dispatcher); + return adapters.get(); + } catch (ReflectiveOperationException error) { + throw new AssertionError("Unable to inspect registered provider adapters", error); + } + } + + private static void assertPrivateFile(Path file) throws IOException { + if (Files.getFileAttributeView(file, PosixFileAttributeView.class) == null) { + return; + } + assertThat(Files.getPosixFilePermissions(file)) + .containsExactlyInAnyOrder(PosixFilePermission.OWNER_READ, PosixFilePermission.OWNER_WRITE); + } + + private static ModelDescriptor existingModel() { + return new ModelDescriptor( + "existing", + "existing-model", + URI.create("https://existing.test/v1"), + ApiStyle.OPENAI_COMPATIBLE, + 0, + 0, + false, + false, + new CostProfile(java.math.BigDecimal.ZERO, java.math.BigDecimal.ZERO, "USD"), + Map.of() + ); + } + + private static void respond(HttpExchange exchange, int status, String body) throws IOException { + byte[] bytes = body.getBytes(StandardCharsets.UTF_8); + exchange.sendResponseHeaders(status, bytes.length); + exchange.getResponseBody().write(bytes); + exchange.close(); + } + + @FunctionalInterface + private interface ExchangeHandler { + void handle(HttpExchange exchange) throws IOException; + } +} diff --git a/lypi-boot/src/test/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfigurationTest.java b/lypi-boot/src/test/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfigurationTest.java index 9b065ac0..38eeb3c2 100644 --- a/lypi-boot/src/test/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfigurationTest.java +++ b/lypi-boot/src/test/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfigurationTest.java @@ -33,6 +33,7 @@ import cn.lypi.contracts.model.AssistantStart; import cn.lypi.contracts.model.AssistantStreamEvent; import cn.lypi.contracts.model.AssistantStreamResult; +import cn.lypi.contracts.model.ModelCatalogPort; import cn.lypi.contracts.model.ModelSelection; import cn.lypi.contracts.model.ThinkingLevel; import cn.lypi.contracts.model.TokenUsage; @@ -53,6 +54,8 @@ import cn.lypi.contracts.runtime.CompactionRuntimePort; import cn.lypi.contracts.runtime.LyPiRuntime; import cn.lypi.contracts.runtime.ResourceRuntimePort; +import cn.lypi.contracts.runtime.ProviderLoginPort; +import cn.lypi.contracts.runtime.ProviderLoginResult; import cn.lypi.contracts.runtime.SecurityRuntimePort; import cn.lypi.contracts.runtime.SessionManagerFactoryPort; import cn.lypi.contracts.runtime.SessionManagerPort; @@ -1597,6 +1600,27 @@ void registersTuiTransportFactoryThatAcceptsSlashCommands() { .run(context -> assertThat(context).hasSingleBean(JLineTuiTransportFactory.class)); } + @Test + void registersTuiTransportFactoryWithModelCatalog() { + ModelCatalogPort catalog = selection -> Optional.empty(); + + new ApplicationContextRunner() + .withUserConfiguration(LyPiRuntimeAutoConfiguration.class) + .withBean(ModelCatalogPort.class, () -> catalog) + .run(context -> assertThat(context).hasSingleBean(JLineTuiTransportFactory.class)); + } + + @Test + void registersTuiTransportFactoryWithProviderLoginPort() { + ProviderLoginPort login = (channelName, baseUrl, authKey) -> + new ProviderLoginResult(channelName, List.of()); + + new ApplicationContextRunner() + .withUserConfiguration(LyPiRuntimeAutoConfiguration.class) + .withBean(ProviderLoginPort.class, () -> login) + .run(context -> assertThat(context).hasSingleBean(JLineTuiTransportFactory.class)); + } + @Test void registersDefaultDiffViewProviderForTuiTransport() { new ApplicationContextRunner() diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/model/ModelCatalogPort.java b/lypi-contracts/src/main/java/cn/lypi/contracts/model/ModelCatalogPort.java index ea41c0f7..2c6550e1 100644 --- a/lypi-contracts/src/main/java/cn/lypi/contracts/model/ModelCatalogPort.java +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/model/ModelCatalogPort.java @@ -1,8 +1,16 @@ package cn.lypi.contracts.model; +import java.util.List; import java.util.Optional; public interface ModelCatalogPort { + /** + * 列出当前进程启动时可用的模型快照。 + */ + default List list() { + return List.of(); + } + /** * 按模型选择查找模型描述。 * diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/runtime/ProviderLoginPort.java b/lypi-contracts/src/main/java/cn/lypi/contracts/runtime/ProviderLoginPort.java new file mode 100644 index 00000000..914bbca1 --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/runtime/ProviderLoginPort.java @@ -0,0 +1,17 @@ +package cn.lypi.contracts.runtime; + +/** + * 注册可在当前进程立即使用的模型 provider。 + * + *

调用方不得记录 auth key,也不得将其写入 session 或事件内容。

+ */ +@FunctionalInterface +public interface ProviderLoginPort { + ProviderLoginResult register(String channelName, String baseUrl, String authKey); + + static ProviderLoginPort unavailable() { + return (channelName, baseUrl, authKey) -> { + throw new IllegalStateException("provider login is unavailable"); + }; + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/runtime/ProviderLoginResult.java b/lypi-contracts/src/main/java/cn/lypi/contracts/runtime/ProviderLoginResult.java new file mode 100644 index 00000000..5398e943 --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/runtime/ProviderLoginResult.java @@ -0,0 +1,15 @@ +package cn.lypi.contracts.runtime; + +import cn.lypi.contracts.model.ModelDescriptor; +import java.util.List; +import java.util.Objects; + +/** + * Provider 登录成功后的非敏感结果。 + */ +public record ProviderLoginResult(String provider, List models) { + public ProviderLoginResult { + provider = Objects.requireNonNull(provider, "provider"); + models = List.copyOf(Objects.requireNonNull(models, "models")); + } +} diff --git a/lypi-contracts/src/test/java/cn/lypi/contracts/CommonContractTest.java b/lypi-contracts/src/test/java/cn/lypi/contracts/CommonContractTest.java index eb9df231..65292d85 100644 --- a/lypi-contracts/src/test/java/cn/lypi/contracts/CommonContractTest.java +++ b/lypi-contracts/src/test/java/cn/lypi/contracts/CommonContractTest.java @@ -2,6 +2,7 @@ import static org.junit.jupiter.api.Assertions.assertAll; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; import cn.lypi.contracts.boundary.BoundaryCheckReport; @@ -29,6 +30,7 @@ import cn.lypi.contracts.runtime.SessionManagerPort; import cn.lypi.contracts.runtime.SessionStorageRootPort; import cn.lypi.contracts.runtime.ToolRuntimePort; +import cn.lypi.contracts.runtime.ProviderLoginPort; import com.fasterxml.jackson.databind.ObjectMapper; import com.fasterxml.jackson.datatype.jdk8.Jdk8Module; import java.lang.reflect.Method; @@ -166,11 +168,22 @@ void runtimePortsExposeDocumentedCrossModuleCapabilities() { () -> assertMethod(ChildSessionPort.class, "create", 1), () -> assertMethod(SessionManagerFactoryPort.class, "open", 2), () -> assertMethod(SessionStorageRootPort.class, "sessionStorageRoot", 0), + () -> assertMethod(ProviderLoginPort.class, "register", 3), () -> assertMethod(ProgressSink.class, "progress", 1), () -> assertMethod(ToolProgressEvent.class, "progress", 0) ); } + @Test + void unavailableProviderLoginUsesTheNamedChannelContract() { + IllegalStateException error = assertThrows( + IllegalStateException.class, + () -> ProviderLoginPort.unavailable().register("zen", "https://example.test/v1", "test-key") + ); + + assertEquals("provider login is unavailable", error.getMessage()); + } + @Test void aiProviderRuntimePortReturnsAssistantEventStream() throws Exception { assertEquals( diff --git a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/JLineTuiTransport.java b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/JLineTuiTransport.java index c6012a49..de290b8f 100644 --- a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/JLineTuiTransport.java +++ b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/JLineTuiTransport.java @@ -7,10 +7,13 @@ import cn.lypi.contracts.event.EventFilter; import cn.lypi.contracts.event.EventSubscription; import cn.lypi.contracts.event.MessageDeltaEvent; +import cn.lypi.contracts.model.ModelCatalogPort; import cn.lypi.contracts.runtime.AgentCorePort; import cn.lypi.contracts.runtime.CompactionRuntimePort; +import cn.lypi.contracts.runtime.ProviderLoginPort; import cn.lypi.contracts.runtime.ResourceRuntimePort; import cn.lypi.contracts.runtime.SessionManagerPort; +import cn.lypi.contracts.session.SessionContext; import cn.lypi.contracts.skill.SkillIndex; import cn.lypi.contracts.tui.DiffViewProvider; import cn.lypi.contracts.tui.NewSessionController; @@ -134,7 +137,8 @@ private JLineTuiTransport( Supplier slashPickerSupplier, DiffViewProvider diffViewProvider, ResumeSessionController resumeController, - Supplier skillIndexSupplier + Supplier skillIndexSupplier, + Supplier modelPickerSupplier ) { this( frameSink, @@ -150,6 +154,7 @@ private JLineTuiTransport( diffViewProvider, resumeController, skillIndexSupplier, + modelPickerSupplier, Clock.systemUTC() ); } @@ -168,6 +173,7 @@ private JLineTuiTransport( DiffViewProvider diffViewProvider, ResumeSessionController resumeController, Supplier skillIndexSupplier, + Supplier modelPickerSupplier, Clock clock ) { this.renderer = null; @@ -188,7 +194,8 @@ private JLineTuiTransport( slashPickerSupplier, resumeController, this::replaceRuntimeState, - skillIndexSupplier + skillIndexSupplier, + modelPickerSupplier ); this.inputPump = new TerminalInputPump(inputSource, new KeyMapper(), inputLoop); this.terminalSession = terminalSession; @@ -266,6 +273,7 @@ public static JLineTuiTransport open( null, diffViewProvider, resumeController, + null, null ); } @@ -425,6 +433,74 @@ public static JLineTuiTransport open( SessionManagerPort sessionManager, ResourceRuntimePort resourceRuntime, CompactionRuntimePort compactionRuntime + ) throws IOException { + return open( + state, + core, + events, + terminal, + diffViewProvider, + slashCommands, + resumeController, + newSessionController, + sessionManager, + resourceRuntime, + compactionRuntime, + null + ); + } + + /** + * 打开真实 JLine TUI transport,并提供只读模型目录。 + */ + public static JLineTuiTransport open( + SessionRuntimeState state, + AgentCorePort core, + EventBus events, + Terminal terminal, + DiffViewProvider diffViewProvider, + List slashCommands, + ResumeSessionController resumeController, + NewSessionController newSessionController, + SessionManagerPort sessionManager, + ResourceRuntimePort resourceRuntime, + CompactionRuntimePort compactionRuntime, + ModelCatalogPort modelCatalog + ) throws IOException { + return open( + state, + core, + events, + terminal, + diffViewProvider, + slashCommands, + resumeController, + newSessionController, + sessionManager, + resourceRuntime, + compactionRuntime, + modelCatalog, + ProviderLoginPort.unavailable() + ); + } + + /** + * 打开真实 JLine TUI transport,并提供只读模型目录和 provider 登录端口。 + */ + public static JLineTuiTransport open( + SessionRuntimeState state, + AgentCorePort core, + EventBus events, + Terminal terminal, + DiffViewProvider diffViewProvider, + List slashCommands, + ResumeSessionController resumeController, + NewSessionController newSessionController, + SessionManagerPort sessionManager, + ResourceRuntimePort resourceRuntime, + CompactionRuntimePort compactionRuntime, + ModelCatalogPort modelCatalog, + ProviderLoginPort providerLogin ) throws IOException { SlashCommandRouter router = new SlashCommandRouter( state.sessionId(), @@ -433,7 +509,8 @@ public static JLineTuiTransport open( resourceRuntime, compactionRuntime, newSessionController, - slashCommands + slashCommands, + modelCatalog ); JLineTuiTransport[] holder = new JLineTuiTransport[1]; RuntimeTuiSubmitHandler submitHandler = new RuntimeTuiSubmitHandler( @@ -446,7 +523,9 @@ public static JLineTuiTransport open( if (holder[0] != null) { holder[0].replaceRuntimeState(runtimeState); } - } + }, + () -> new SkillIndex(List.of(), List.of()), + providerLogin ); JLineTuiTransport transport = openTerminal( state, @@ -456,12 +535,24 @@ public static JLineTuiTransport open( () -> new SlashCommandPicker(router.commandNames()), diffViewProvider, resumeController, - () -> resourceRuntime.load(state.cwd()).skillIndex() + () -> resourceRuntime.load(state.cwd()).skillIndex(), + modelPickerSupplier(modelCatalog, router, state) ); holder[0] = transport; return transport; } + private static Supplier modelPickerSupplier( + ModelCatalogPort modelCatalog, + SlashCommandRouter router, + SessionRuntimeState state + ) { + return () -> new ModelPicker( + modelCatalog == null ? List.of() : Optional.ofNullable(modelCatalog.list()).orElse(List.of()), + router.sessionContext().map(SessionContext::model).orElse(state.model()) + ); + } + public static JLineTuiTransport open( SessionRuntimeState state, AgentCorePort core, @@ -551,6 +642,80 @@ static JLineTuiTransport open( NewSessionController newSessionController, int width, int height + ) throws IOException { + return open( + state, + core, + events, + io, + inputSource, + slashCommands, + sessionManager, + resourceRuntime, + compactionRuntime, + diffViewProvider, + resumeController, + newSessionController, + null, + width, + height + ); + } + + static JLineTuiTransport open( + SessionRuntimeState state, + AgentCorePort core, + EventBus events, + TerminalIo io, + TerminalInputSource inputSource, + List slashCommands, + SessionManagerPort sessionManager, + ResourceRuntimePort resourceRuntime, + CompactionRuntimePort compactionRuntime, + DiffViewProvider diffViewProvider, + ResumeSessionController resumeController, + NewSessionController newSessionController, + ModelCatalogPort modelCatalog, + int width, + int height + ) throws IOException { + return open( + state, + core, + events, + io, + inputSource, + slashCommands, + sessionManager, + resourceRuntime, + compactionRuntime, + diffViewProvider, + resumeController, + newSessionController, + modelCatalog, + ProviderLoginPort.unavailable(), + width, + height + ); + } + + static JLineTuiTransport open( + SessionRuntimeState state, + AgentCorePort core, + EventBus events, + TerminalIo io, + TerminalInputSource inputSource, + List slashCommands, + SessionManagerPort sessionManager, + ResourceRuntimePort resourceRuntime, + CompactionRuntimePort compactionRuntime, + DiffViewProvider diffViewProvider, + ResumeSessionController resumeController, + NewSessionController newSessionController, + ModelCatalogPort modelCatalog, + ProviderLoginPort providerLogin, + int width, + int height ) throws IOException { SlashCommandRouter router = new SlashCommandRouter( state.sessionId(), @@ -559,7 +724,8 @@ static JLineTuiTransport open( resourceRuntime, compactionRuntime, newSessionController, - slashCommands + slashCommands, + modelCatalog ); JLineTuiTransport[] holder = new JLineTuiTransport[1]; RuntimeTuiSubmitHandler submitHandler = new RuntimeTuiSubmitHandler( @@ -572,7 +738,9 @@ static JLineTuiTransport open( if (holder[0] != null) { holder[0].replaceRuntimeState(runtimeState); } - } + }, + () -> new SkillIndex(List.of(), List.of()), + providerLogin ); JLineTuiTransport transport = open( state, @@ -584,6 +752,7 @@ static JLineTuiTransport open( diffViewProvider, resumeController, () -> resourceRuntime.load(state.cwd()).skillIndex(), + modelPickerSupplier(modelCatalog, router, state), width, height ); @@ -669,6 +838,7 @@ static JLineTuiTransport withBatchInput( null, NOOP_DIFF_VIEW_PROVIDER, null, + null, null ); } @@ -713,6 +883,7 @@ static JLineTuiTransport withBatchInput( NOOP_DIFF_VIEW_PROVIDER, null, null, + null, clock ); } @@ -736,7 +907,8 @@ private static JLineTuiTransport openTerminal( Supplier slashPickerSupplier, DiffViewProvider diffViewProvider, ResumeSessionController resumeController, - Supplier skillIndexSupplier + Supplier skillIndexSupplier, + Supplier modelPickerSupplier ) throws IOException { JLineTerminalIo io = new JLineTerminalIo(terminal); JLineTuiTransport[] holder = new JLineTuiTransport[1]; @@ -780,7 +952,8 @@ private static JLineTuiTransport openTerminal( slashPickerSupplier, diffViewProvider, resumeController, - skillIndexSupplier + skillIndexSupplier, + modelPickerSupplier ); holder[0] = transport; transport.attach(events, state); @@ -860,6 +1033,36 @@ static JLineTuiTransport open( Supplier skillIndexSupplier, int width, int height + ) throws IOException { + return open( + state, + events, + io, + inputSource, + submitHandler, + slashPickerSupplier, + diffViewProvider, + resumeController, + skillIndexSupplier, + null, + width, + height + ); + } + + static JLineTuiTransport open( + SessionRuntimeState state, + EventBus events, + TerminalIo io, + TerminalInputSource inputSource, + TuiSubmitHandler submitHandler, + Supplier slashPickerSupplier, + DiffViewProvider diffViewProvider, + ResumeSessionController resumeController, + Supplier skillIndexSupplier, + Supplier modelPickerSupplier, + int width, + int height ) throws IOException { JLineTuiTransport[] holder = new JLineTuiTransport[1]; TerminalSession session = TerminalSession.open(io, () -> { @@ -894,7 +1097,8 @@ static JLineTuiTransport open( slashPickerSupplier, diffViewProvider, resumeController, - skillIndexSupplier + skillIndexSupplier, + modelPickerSupplier ); holder[0] = transport; transport.attach(events, state); diff --git a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/LoginOverlay.java b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/LoginOverlay.java new file mode 100644 index 00000000..981caea8 --- /dev/null +++ b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/LoginOverlay.java @@ -0,0 +1,107 @@ +package cn.lypi.transport.tui; + +import java.util.List; +import java.util.Optional; + +/** Temporary three-step input state for provider login credentials. */ +final class LoginOverlay { + private enum Step { + CHANNEL_NAME, + BASE_URL, + AUTH_KEY + } + + private final StringBuilder channelName = new StringBuilder(); + private final StringBuilder baseUrl = new StringBuilder(); + private final StringBuilder authKey = new StringBuilder(); + private Step step = Step.CHANNEL_NAME; + private boolean open; + + void open() { + clear(); + open = true; + } + + void append(String text) { + if (text != null) { + open = true; + currentInput().append(text); + } + } + + void backspace() { + StringBuilder input = currentInput(); + if (!input.isEmpty()) { + input.deleteCharAt(input.length() - 1); + } + } + + Optional accept() { + if (step == Step.CHANNEL_NAME) { + if (channelName.toString().isBlank()) { + return Optional.empty(); + } + step = Step.BASE_URL; + return Optional.empty(); + } + if (step == Step.BASE_URL) { + if (baseUrl.toString().isBlank()) { + return Optional.empty(); + } + step = Step.AUTH_KEY; + return Optional.empty(); + } + if (authKey.toString().isBlank()) { + return Optional.empty(); + } + return Optional.of(new Submission(channelName.toString(), baseUrl.toString(), authKey.toString())); + } + + void clear() { + clear(channelName); + clear(baseUrl); + clear(authKey); + step = Step.CHANNEL_NAME; + open = false; + } + + List lines() { + if (!open) { + return List.of(); + } + return switch (step) { + case CHANNEL_NAME -> List.of("Channel name: " + channelName); + case BASE_URL -> List.of( + "Channel name: " + channelName, + "Base URL: " + baseUrl + ); + case AUTH_KEY -> List.of( + "Channel name: " + channelName, + "Base URL: " + baseUrl, + "Auth key: " + "*".repeat(authKey.length()) + ); + }; + } + + private StringBuilder currentInput() { + return switch (step) { + case CHANNEL_NAME -> channelName; + case BASE_URL -> baseUrl; + case AUTH_KEY -> authKey; + }; + } + + private static void clear(StringBuilder value) { + for (int index = 0; index < value.length(); index++) { + value.setCharAt(index, '\0'); + } + value.setLength(0); + } + + record Submission(String channelName, String baseUrl, String authKey) { + @Override + public String toString() { + return "Submission[channelName=" + channelName + ", baseUrl=" + baseUrl + ", authKey=]"; + } + } +} diff --git a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/ModelPicker.java b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/ModelPicker.java new file mode 100644 index 00000000..d621dfff --- /dev/null +++ b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/ModelPicker.java @@ -0,0 +1,79 @@ +package cn.lypi.transport.tui; + +import cn.lypi.contracts.model.ModelDescriptor; +import cn.lypi.contracts.model.ModelSelection; +import java.util.Comparator; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; + +final class ModelPicker { + private final List models; + private int selectedIndex; + + ModelPicker(List models, ModelSelection current) { + this.models = distinctAndSort(models); + this.selectedIndex = initialIndex(current); + } + + List labels() { + return models.stream().map(ModelPicker::label).toList(); + } + + int selectedIndex() { + return selectedIndex; + } + + Optional selectedLabel() { + return accept().map(ModelPicker::label); + } + + Optional accept() { + return models.isEmpty() ? Optional.empty() : Optional.of(models.get(selectedIndex)); + } + + void moveDown() { + if (!models.isEmpty()) { + selectedIndex = Math.floorMod(selectedIndex + 1, models.size()); + } + } + + void moveUp() { + if (!models.isEmpty()) { + selectedIndex = Math.floorMod(selectedIndex - 1, models.size()); + } + } + + static String label(ModelDescriptor model) { + return model.provider() + "/" + model.modelId(); + } + + private static List distinctAndSort(List candidates) { + Map distinct = new LinkedHashMap<>(); + for (ModelDescriptor candidate : candidates == null ? List.of() : candidates) { + if (candidate != null) { + distinct.putIfAbsent(new ModelKey(candidate.provider(), candidate.modelId()), candidate); + } + } + return distinct.values().stream() + .sorted(Comparator.comparing(ModelDescriptor::provider).thenComparing(ModelDescriptor::modelId)) + .toList(); + } + + private int initialIndex(ModelSelection current) { + if (current == null) { + return 0; + } + for (int index = 0; index < models.size(); index++) { + ModelDescriptor candidate = models.get(index); + if (candidate.provider().equals(current.provider()) && candidate.modelId().equals(current.modelId())) { + return index; + } + } + return 0; + } + + private record ModelKey(String provider, String modelId) { + } +} diff --git a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/RuntimeTuiSubmitHandler.java b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/RuntimeTuiSubmitHandler.java index 2982483f..24766ca9 100644 --- a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/RuntimeTuiSubmitHandler.java +++ b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/RuntimeTuiSubmitHandler.java @@ -18,9 +18,12 @@ import cn.lypi.contracts.event.MessageStartEvent; import cn.lypi.contracts.event.PermissionResponseEvent; import cn.lypi.contracts.event.SessionStateEvent; +import cn.lypi.contracts.error.ModelProviderException; import cn.lypi.contracts.session.SessionContext; import cn.lypi.contracts.runtime.AgentCorePort; import cn.lypi.contracts.runtime.CompactionResult; +import cn.lypi.contracts.runtime.ProviderLoginPort; +import cn.lypi.contracts.runtime.ProviderLoginResult; import cn.lypi.contracts.skill.SkillIndex; import cn.lypi.contracts.skill.SkillMention; import cn.lypi.contracts.tui.SessionRuntimeState; @@ -47,10 +50,12 @@ final class RuntimeTuiSubmitHandler implements TuiSubmitHandler { private final SlashCommandRouter slashCommandRouter; private final Consumer runtimeStateConsumer; private final Supplier skillIndexSupplier; + private final ProviderLoginPort providerLogin; private final Object activeTurnLock = new Object(); private volatile MutableAbortSignal activeSignal; private ActiveTurn activeTurn; private volatile boolean compactRunning; + private final AtomicBoolean providerLoginRunning = new AtomicBoolean(); RuntimeTuiSubmitHandler(String sessionId, AgentCorePort core, EventBus events) { this(sessionId, core, events, command -> Thread.ofVirtual().name("lypi-tui-turn-", 0).start(command)); @@ -117,6 +122,28 @@ final class RuntimeTuiSubmitHandler implements TuiSubmitHandler { SlashCommandRouter slashCommandRouter, Consumer runtimeStateConsumer, Supplier skillIndexSupplier + ) { + this( + sessionId, + core, + events, + executor, + slashCommandRouter, + runtimeStateConsumer, + skillIndexSupplier, + ProviderLoginPort.unavailable() + ); + } + + RuntimeTuiSubmitHandler( + String sessionId, + AgentCorePort core, + EventBus events, + Executor executor, + SlashCommandRouter slashCommandRouter, + Consumer runtimeStateConsumer, + Supplier skillIndexSupplier, + ProviderLoginPort providerLogin ) { this.currentSessionId = sessionId; this.core = core; @@ -127,6 +154,7 @@ final class RuntimeTuiSubmitHandler implements TuiSubmitHandler { this.skillIndexSupplier = skillIndexSupplier == null ? () -> new SkillIndex(List.of(), List.of()) : skillIndexSupplier; + this.providerLogin = providerLogin == null ? ProviderLoginPort.unavailable() : providerLogin; } @Override @@ -134,6 +162,35 @@ public void submitUserInput(String input) { submitUserInput(input, List.of()); } + @Override + public void submitProviderLogin(String channelName, String baseUrl, String authKey) { + if (!providerLoginRunning.compareAndSet(false, true)) { + publishSlashCommandError("login: provider registration is running"); + return; + } + try { + executor.execute(() -> runProviderLogin(channelName, baseUrl, authKey)); + } catch (RuntimeException error) { + providerLoginRunning.set(false); + publishSlashCommandError("login: provider registration failed"); + } + } + + private void runProviderLogin(String channelName, String baseUrl, String authKey) { + try { + ProviderLoginResult result = providerLogin.register(channelName, baseUrl, authKey); + int modelCount = result.models().size(); + publishSlashCommandNotice( + "login: registered " + result.provider() + " (" + modelCount + " model" + + (modelCount == 1 ? "" : "s") + ")" + ); + } catch (RuntimeException error) { + publishSlashCommandError("login: " + safeLoginError(error, authKey)); + } finally { + providerLoginRunning.set(false); + } + } + @Override public void submitUserInput(String input, List skillMentions) { String routedInput = input == null ? "" : input; @@ -435,6 +492,20 @@ private String errorMessage(RuntimeException exception) { return message == null || message.isBlank() ? exception.getClass().getSimpleName() : message; } + private static String safeLoginError(RuntimeException exception, String authKey) { + if (exception instanceof ModelProviderException) { + String message = exception.getMessage(); + if (message != null && !message.isBlank() && (authKey == null || !message.contains(authKey))) { + return message; + } + } + if (exception instanceof IllegalStateException + && "provider login is unavailable".equals(exception.getMessage())) { + return "provider login is unavailable"; + } + return "provider registration failed"; + } + private void publishSessionState() { slashCommandRouter.sessionContext().ifPresent(context -> events.publish(new SessionStateEvent( currentSessionId, diff --git a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/SlashCommandPicker.java b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/SlashCommandPicker.java index d0fc2cda..5108f6cf 100644 --- a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/SlashCommandPicker.java +++ b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/SlashCommandPicker.java @@ -7,6 +7,7 @@ final class SlashCommandPicker { private static final List BUILT_IN_COMMANDS = List.of( "/model", + "/login", "/thinking", "/plan", "/permission-mode", diff --git a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/SlashCommandRouter.java b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/SlashCommandRouter.java index a77907a0..cc9686d4 100644 --- a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/SlashCommandRouter.java +++ b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/SlashCommandRouter.java @@ -1,8 +1,10 @@ package cn.lypi.transport.tui; +import cn.lypi.contracts.common.AbortSignal; +import cn.lypi.contracts.model.ModelCatalogPort; +import cn.lypi.contracts.model.ModelDescriptor; import cn.lypi.contracts.model.ModelSelection; import cn.lypi.contracts.model.ThinkingLevel; -import cn.lypi.contracts.common.AbortSignal; import cn.lypi.contracts.prompt.PromptParameter; import cn.lypi.contracts.prompt.PromptRenderRequest; import cn.lypi.contracts.prompt.PromptRenderResult; @@ -36,6 +38,7 @@ final class SlashCommandRouter { private static final List BUILT_IN_COMMANDS = List.of( "/compact", + "/login", "/model", "/new", "/permission-mode", @@ -50,6 +53,7 @@ final class SlashCommandRouter { private final CompactionRuntimePort compactionRuntime; private final NewSessionController newSessionController; private final List slashCommands; + private final ModelCatalogPort modelCatalog; SlashCommandRouter( String sessionId, @@ -89,6 +93,28 @@ final class SlashCommandRouter { CompactionRuntimePort compactionRuntime, NewSessionController newSessionController, List slashCommands + ) { + this( + sessionId, + cwd, + sessionManager, + resourceRuntime, + compactionRuntime, + newSessionController, + slashCommands, + null + ); + } + + SlashCommandRouter( + String sessionId, + Path cwd, + SessionManagerPort sessionManager, + ResourceRuntimePort resourceRuntime, + CompactionRuntimePort compactionRuntime, + NewSessionController newSessionController, + List slashCommands, + ModelCatalogPort modelCatalog ) { this.sessionId = Objects.requireNonNull(sessionId, "sessionId must not be null"); this.cwd = cwd == null ? Path.of(".") : cwd; @@ -97,6 +123,7 @@ final class SlashCommandRouter { this.compactionRuntime = compactionRuntime; this.newSessionController = newSessionController; this.slashCommands = safeSlashCommands(slashCommands); + this.modelCatalog = modelCatalog; } SlashCommandRouter(List slashCommands) { @@ -107,6 +134,7 @@ final class SlashCommandRouter { this.compactionRuntime = null; this.newSessionController = null; this.slashCommands = safeSlashCommands(slashCommands); + this.modelCatalog = null; } SlashCommandResult route(String input) { @@ -123,6 +151,7 @@ SlashCommandResult route(String input) { return SlashCommandResult.notMatched(); } return switch (match.command().orElseThrow()) { + case "/login" -> routeLogin(arguments); case "/model" -> routeModel(arguments, input); case "/thinking" -> routeThinking(arguments, input); case "/permission-mode" -> routePermissionMode(arguments, input); @@ -377,16 +406,35 @@ private SlashCommandResult routeModel(SlashCommandArguments arguments, String re provider = modelId.substring(0, separator); modelId = modelId.substring(separator + 1); } + ThinkingLevel thinkingLevel = context.thinkingLevel(); + if (modelCatalog != null) { + ModelSelection lookup = new ModelSelection(provider, modelId, thinkingLevel); + Optional descriptor = modelCatalog.find(lookup); + if (descriptor.isEmpty()) { + return SlashCommandResult.error("unknown model: " + provider + "/" + modelId); + } + if (!descriptor.orElseThrow().supportsThinking()) { + thinkingLevel = ThinkingLevel.OFF; + } + } + ModelSelection selection = new ModelSelection(provider, modelId, thinkingLevel); append(new ModelChangeEntry( newEntryId(), leafId, - new ModelSelection(provider, modelId, context.thinkingLevel()), + selection, reason, Instant.now() )); return SlashCommandResult.stateChangedNotice("model: " + provider + "/" + modelId); } + private SlashCommandResult routeLogin(SlashCommandArguments arguments) { + if (!arguments.positionals().isEmpty() || !arguments.named().isEmpty()) { + return SlashCommandResult.error("usage: /login"); + } + return SlashCommandResult.consumedCommand(); + } + private SlashCommandResult routeThinking(SlashCommandArguments arguments, String reason) { if (arguments.positionals().size() != 1) { return SlashCommandResult.error("usage: /thinking "); diff --git a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/TuiInputLoop.java b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/TuiInputLoop.java index fdc6c1c2..9a8b9d5e 100644 --- a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/TuiInputLoop.java +++ b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/TuiInputLoop.java @@ -1,6 +1,7 @@ package cn.lypi.transport.tui; import cn.lypi.contracts.agent.SteeringMessage; +import cn.lypi.contracts.model.ModelDescriptor; import cn.lypi.contracts.tui.TuiBlock; import cn.lypi.contracts.tui.TuiMessageBlock; import cn.lypi.contracts.tui.PermissionPromptView; @@ -25,16 +26,21 @@ final class TuiInputLoop { private final KeyBindingRegistry bindings = KeyBindingRegistry.defaults(); private final TerminalInputPolicy inputPolicy = new TerminalInputPolicy(); private final Supplier slashPickerSupplier; + private final Supplier modelPickerSupplier; private final Supplier skillIndexSupplier; private final Runnable renderRequest; private final ResumeSessionController resumeController; private final ResumeOverlayController resumeOverlayController; private SlashCommandPicker slashPicker; + private ModelPicker modelPicker; private SkillMentionToken skillToken; private int skillSelectedIndex; private final List skillBindings = new java.util.ArrayList<>(); private final SkillMentionSuppressions skillSuppressions = new SkillMentionSuppressions(); private boolean slashOverlayClosed; + private boolean modelOverlayOpen; + private final LoginOverlay loginOverlay = new LoginOverlay(); + private boolean loginOverlayOpen; private boolean interruptibleRunning; private boolean exitRequested; private boolean toolOutputExpanded; @@ -100,6 +106,30 @@ final class TuiInputLoop { ResumeSessionController resumeController, Consumer resumeStateConsumer, Supplier skillIndexSupplier + ) { + this( + submitHandler, + renderRequest, + layout, + viewSupplier, + slashPickerSupplier, + resumeController, + resumeStateConsumer, + skillIndexSupplier, + null + ); + } + + TuiInputLoop( + TuiSubmitHandler submitHandler, + Runnable renderRequest, + TuiLayout layout, + Supplier viewSupplier, + Supplier slashPickerSupplier, + ResumeSessionController resumeController, + Consumer resumeStateConsumer, + Supplier skillIndexSupplier, + Supplier modelPickerSupplier ) { this.submitHandler = submitHandler; this.layout = layout; @@ -107,6 +137,9 @@ final class TuiInputLoop { this.slashPickerSupplier = slashPickerSupplier == null ? () -> SlashCommandPicker.withTemplates(List.of()) : slashPickerSupplier; + this.modelPickerSupplier = modelPickerSupplier == null + ? () -> new ModelPicker(List.of(), null) + : modelPickerSupplier; this.skillIndexSupplier = skillIndexSupplier == null ? () -> new SkillIndex(List.of(), List.of()) : skillIndexSupplier; this.renderRequest = renderRequest == null ? () -> { } : renderRequest; @@ -128,6 +161,15 @@ void acceptText(String text) { render(); return; } + if (loginOverlayOpen) { + loginOverlay.append(text); + render(); + return; + } + if (modelOverlayOpen) { + render(); + return; + } if (resumeOverlayController != null) { resumeOverlayController.clearTransientLine(); } @@ -142,6 +184,15 @@ void acceptPaste(String text) { render(); return; } + if (loginOverlayOpen) { + loginOverlay.append(text); + render(); + return; + } + if (modelOverlayOpen) { + render(); + return; + } if (resumeOverlayController != null) { resumeOverlayController.clearTransientLine(); } @@ -182,6 +233,10 @@ void acceptKey(TerminalKey key) { return; } } + if (loginOverlayOpen && !resumeOverlayOpen()) { + handleLoginOverlayKey(key); + return; + } if ((key == TerminalKey.ESC || key == TerminalKey.CTRL_C) && interruptibleRunning && submitHandler.hasPendingSteeringMessages()) { @@ -206,6 +261,10 @@ void acceptKey(TerminalKey key) { handleResumeOverlayKey(key); return; } + if (modelOverlayOpen) { + handleModelOverlayKey(key); + return; + } if (key == TerminalKey.ENTER && "/resume".equals(editor.text().trim()) && resumeController != null) { submitDraft(); return; @@ -342,6 +401,23 @@ private void submitDraft() { render(); return; } + if ("/model".equals(draft.trim())) { + editor.clear(); + slashOverlayClosed = true; + skillBindings.clear(); + skillSuppressions.clear(); + openModelOverlay(); + render(); + return; + } + if ("/login".equals(draft.trim())) { + editor.clear(); + skillBindings.clear(); + skillSuppressions.clear(); + openLoginOverlay(); + render(); + return; + } editor.acceptHistoryEntry(); slashOverlayClosed = true; List mentions = new SkillMentionParser(skillIndexSupplier.get().skills()) @@ -480,6 +556,8 @@ private boolean hasOptionId(PermissionPromptView prompt, String optionId) { private boolean slashOverlayOpen() { return viewSupplier.get().permissionPrompt().isEmpty() && !resumeOverlayOpen() + && !loginOverlayOpen + && !modelOverlayOpen && !slashOverlayClosed && slashFilter().isPresent(); } @@ -547,6 +625,14 @@ List overlayLines() { return resumeLines; } } + List loginLines = loginOverlayLines(); + if (!loginLines.isEmpty()) { + return loginLines; + } + List modelLines = modelOverlayLines(); + if (!modelLines.isEmpty()) { + return modelLines; + } List skillLines = skillOverlayLines(); if (!skillLines.isEmpty()) { return skillLines; @@ -567,13 +653,144 @@ private void handleResumeOverlayKey(TerminalKey key) { } private void acceptSlashSelection() { - slashPicker().accept().ifPresent(command -> editor.replaceFirstToken(command + " ")); + Optional selected = slashPicker().accept(); + if (selected.isPresent() && "/model".equals(selected.orElseThrow())) { + editor.clear(); + skillBindings.clear(); + skillSuppressions.clear(); + openModelOverlay(); + } else if (selected.isPresent() && "/login".equals(selected.orElseThrow())) { + skillBindings.clear(); + skillSuppressions.clear(); + openLoginOverlay(); + } else { + selected.ifPresent(command -> editor.replaceFirstToken(command + " ")); + } slashOverlayClosed = true; render(); } + private void openModelOverlay() { + ModelPicker next = modelPickerSupplier.get(); + modelPicker = next == null ? new ModelPicker(List.of(), null) : next; + modelOverlayOpen = true; + } + + private void closeModelOverlay() { + modelOverlayOpen = false; + modelPicker = null; + } + + private void openLoginOverlay() { + editor.clear(); + closeModelOverlay(); + loginOverlay.open(); + loginOverlayOpen = true; + slashOverlayClosed = true; + skillToken = null; + } + + private void closeLoginOverlay() { + loginOverlay.clear(); + loginOverlayOpen = false; + } + + private void handleLoginOverlayKey(TerminalKey key) { + if (key == TerminalKey.ESC || key == TerminalKey.CTRL_C) { + closeLoginOverlay(); + render(); + return; + } + if (key == TerminalKey.BACKSPACE) { + loginOverlay.backspace(); + render(); + return; + } + if (key == TerminalKey.ENTER) { + Optional submission = loginOverlay.accept(); + if (submission.isPresent()) { + LoginOverlay.Submission value = submission.orElseThrow(); + closeLoginOverlay(); + submitHandler.submitProviderLogin(value.channelName(), value.baseUrl(), value.authKey()); + } + render(); + return; + } + render(); + } + + private void handleModelOverlayKey(TerminalKey key) { + if (key == TerminalKey.ESC) { + closeModelOverlay(); + render(); + return; + } + if (key == TerminalKey.UP) { + modelPicker().moveUp(); + render(); + return; + } + if (key == TerminalKey.DOWN) { + modelPicker().moveDown(); + render(); + return; + } + if (key == TerminalKey.ENTER) { + Optional selected = modelPicker().accept(); + if (selected.isPresent()) { + submitHandler.submitUserInput("/model " + ModelPicker.label(selected.orElseThrow())); + editor.clear(); + closeModelOverlay(); + } + render(); + return; + } + render(); + } + + private ModelPicker modelPicker() { + if (modelPicker == null) { + openModelOverlay(); + } + return modelPicker; + } + + private List modelOverlayLines() { + if (!modelOverlayVisible()) { + return List.of(); + } + ModelPicker picker = modelPicker(); + List labels = picker.labels(); + if (labels.isEmpty()) { + return List.of("No models available"); + } + int limit = Math.min(labels.size(), Math.max(1, Math.min(10, layout.height() - 4))); + int selected = Math.max(0, Math.min(picker.selectedIndex(), labels.size() - 1)); + int start = Math.max(0, Math.min(selected, labels.size() - limit)); + List lines = new ArrayList<>(limit); + for (int index = start; index < start + limit; index++) { + lines.add((index == selected ? "> " : " ") + labels.get(index)); + } + return List.copyOf(lines); + } + + private boolean modelOverlayVisible() { + return modelOverlayOpen + && viewSupplier.get().permissionPrompt().isEmpty() + && !resumeOverlayOpen() + && !loginOverlayOpen; + } + + private List loginOverlayLines() { + return loginOverlayOpen ? loginOverlay.lines() : List.of(); + } + private boolean skillOverlayOpen() { - if (viewSupplier.get().permissionPrompt().isPresent() || resumeOverlayOpen() || slashOverlayOpen()) { + if (viewSupplier.get().permissionPrompt().isPresent() + || resumeOverlayOpen() + || loginOverlayOpen + || modelOverlayOpen + || slashOverlayOpen()) { return false; } SkillMentionParser parser = new SkillMentionParser(skillIndexSupplier.get().skills()); diff --git a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/TuiSubmitHandler.java b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/TuiSubmitHandler.java index 32e3812b..4ddae1ec 100644 --- a/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/TuiSubmitHandler.java +++ b/lypi-transport-tui/src/main/java/cn/lypi/transport/tui/TuiSubmitHandler.java @@ -18,6 +18,12 @@ default void submitUserInput(String input, List skillMentions) { submitUserInput(input); } + /** + * Submits temporary provider-login credentials without creating a user turn. + */ + default void submitProviderLogin(String channelName, String baseUrl, String authKey) { + } + default List pendingSteeringMessages() { return List.of(); } diff --git a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/JLineTuiTransportTest.java b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/JLineTuiTransportTest.java index 8fa4bef5..d366aed4 100644 --- a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/JLineTuiTransportTest.java +++ b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/JLineTuiTransportTest.java @@ -22,15 +22,23 @@ import cn.lypi.contracts.event.EventSubscription; import cn.lypi.contracts.event.ErrorEvent; import cn.lypi.contracts.event.MessageDeltaEvent; +import cn.lypi.contracts.event.SessionStateEvent; +import cn.lypi.contracts.model.ApiStyle; +import cn.lypi.contracts.model.CostProfile; +import cn.lypi.contracts.model.ModelCatalogPort; +import cn.lypi.contracts.model.ModelDescriptor; import cn.lypi.contracts.model.ModelSelection; import cn.lypi.contracts.model.ThinkingLevel; import cn.lypi.contracts.resource.ResourceSnapshot; import cn.lypi.contracts.runtime.AgentCorePort; +import cn.lypi.contracts.runtime.ProviderLoginPort; +import cn.lypi.contracts.runtime.ProviderLoginResult; import cn.lypi.contracts.runtime.ResourceRuntimePort; import cn.lypi.contracts.runtime.SessionManagerPort; import cn.lypi.contracts.security.AgentMode; import cn.lypi.contracts.security.PermissionMode; import cn.lypi.contracts.session.ForkRequest; +import cn.lypi.contracts.session.ModelChangeEntry; import cn.lypi.contracts.session.SessionContext; import cn.lypi.contracts.session.SessionEntry; import cn.lypi.contracts.session.SessionHandle; @@ -53,6 +61,7 @@ import java.io.StringWriter; import java.lang.reflect.Proxy; import java.math.BigDecimal; +import java.net.URI; import java.nio.file.Path; import java.time.Duration; import java.time.Instant; @@ -61,6 +70,8 @@ import java.util.List; import java.util.Map; import java.util.Optional; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; import org.jline.terminal.Attributes; import org.jline.terminal.Terminal; import org.jline.utils.NonBlockingReader; @@ -142,6 +153,55 @@ void openAssemblesTerminalSessionRendererInputAndEventSubscription() throws Exce )); } + @Test + void loginOverlaySubmitsCredentialsThroughInjectedProviderPort() throws Exception { + String authKey = "test-secret"; + RecordingTerminalIo io = new RecordingTerminalIo(); + RecordingEventBus events = new RecordingEventBus(); + RecordingCore core = new RecordingCore(); + RecordingSessionManager session = new RecordingSessionManager(); + RecordingProviderLoginPort login = new RecordingProviderLoginPort( + "zen", + "https://api.example.test/v1", + authKey + ); + JLineTuiTransport transport = JLineTuiTransport.open( + runtimeState(), + core, + events, + io, + new QueueInputSource( + "/login", "\r", + "zen", "\r", + "https://api.example.test/v1", "\r", + authKey, "\r" + ), + List.of(), + session, + emptyResources(), + null, + NOOP_DIFF_PROVIDER, + null, + null, + null, + login, + 80, + 8 + ); + + transport.drainInputForTest(); + + assertTrue(login.registered.await(2, TimeUnit.SECONDS)); + assertTrue(login.acceptedChannelName); + assertTrue(login.acceptedBaseUrl); + assertTrue(login.acceptedAuthKey); + assertTrue(core.requests.isEmpty()); + assertTrue(session.entries.isEmpty()); + assertFalse(io.output.toString().contains(authKey)); + + transport.close(); + } + @Test void publicOpenUsesCursorProbeAnchorAndReplaysConcurrentInput() throws Exception { StringWriter output = new StringWriter(); @@ -597,6 +657,54 @@ void openWithRuntimePortsRoutesSlashCommandsBeforeCoreSubmission() throws Except transport.close(); } + @Test + void openWithModelCatalogSelectsModelAndPublishesSessionState() throws Exception { + RecordingTerminalIo io = new RecordingTerminalIo(); + io.width = 80; + io.height = 10; + RecordingEventBus events = new RecordingEventBus(); + RecordingCore core = new RecordingCore(); + RecordingSessionManager session = new RecordingSessionManager(); + ModelCatalogPort catalog = modelCatalog(List.of( + model("openai", "gpt-5"), + model("zen", "kimi-k2.6") + )); + + JLineTuiTransport transport = JLineTuiTransport.open( + runtimeState(), + core, + events, + io, + new QueueInputSource("/model", "\r", "\033[B", "\r"), + List.of(), + session, + emptyResources(), + null, + NOOP_DIFF_PROVIDER, + null, + null, + catalog, + 80, + 10 + ); + + transport.drainInputForTest(); + + ModelChangeEntry entry = assertInstanceOf(ModelChangeEntry.class, session.entries.getFirst()); + assertEquals(new ModelSelection("zen", "kimi-k2.6", ThinkingLevel.MEDIUM), entry.model()); + SessionStateEvent stateEvent = events.published.stream() + .filter(SessionStateEvent.class::isInstance) + .map(SessionStateEvent.class::cast) + .findFirst() + .orElseThrow(); + assertEquals(entry.model(), stateEvent.model()); + assertEquals(0, core.requests.size()); + assertTrue(io.output.toString().contains("openai/gpt-5")); + assertTrue(io.output.toString().contains("zen/kimi-k2.6")); + + transport.close(); + } + @Test void resumeRuntimeStateRebindsEventSubscriptionToResumedSession() throws Exception { RecordingTerminalIo io = new RecordingTerminalIo(); @@ -946,9 +1054,43 @@ public cn.lypi.contracts.prompt.SystemPrompt buildSystemPrompt(ResourceSnapshot }; } + private static ModelCatalogPort modelCatalog(List descriptors) { + List models = List.copyOf(descriptors); + return new ModelCatalogPort() { + @Override + public List list() { + return models; + } + + @Override + public Optional find(ModelSelection selection) { + return models.stream() + .filter(candidate -> candidate.provider().equals(selection.provider())) + .filter(candidate -> candidate.modelId().equals(selection.modelId())) + .findFirst(); + } + }; + } + + private static ModelDescriptor model(String provider, String modelId) { + return new ModelDescriptor( + provider, + modelId, + URI.create("https://api.example.test/v1"), + ApiStyle.OPENAI_COMPATIBLE, + 128_000, + 16_384, + true, + false, + new CostProfile(BigDecimal.ZERO, BigDecimal.ZERO, "USD"), + Map.of() + ); + } + private static final class RecordingSessionManager implements SessionManagerPort { private final List entries = new ArrayList<>(); private String leafId = "root"; + private ModelSelection model = new ModelSelection("openai", "gpt-5", ThinkingLevel.MEDIUM); @Override public SessionHandle openOrCreate(String sessionId) { @@ -958,6 +1100,9 @@ public SessionHandle openOrCreate(String sessionId) { @Override public SessionHandle append(SessionEntry entry) { entries.add(entry); + if (entry instanceof ModelChangeEntry modelChange) { + model = modelChange.model(); + } leafId = entry.id(); return openOrCreate("ses_1"); } @@ -994,7 +1139,7 @@ public SessionContext context(String leafId) { List.of(), List.of(this.leafId), List.of(), - new ModelSelection("openai", "gpt-5", ThinkingLevel.MEDIUM), + model, ThinkingLevel.MEDIUM, AgentMode.EXECUTE, PermissionMode.ASK @@ -1022,6 +1167,31 @@ public void requestInterrupt(String reason) { } } + private static final class RecordingProviderLoginPort implements ProviderLoginPort { + private final String expectedChannelName; + private final String expectedBaseUrl; + private final String expectedAuthKey; + private final CountDownLatch registered = new CountDownLatch(1); + private volatile boolean acceptedChannelName; + private volatile boolean acceptedBaseUrl; + private volatile boolean acceptedAuthKey; + + private RecordingProviderLoginPort(String expectedChannelName, String expectedBaseUrl, String expectedAuthKey) { + this.expectedChannelName = expectedChannelName; + this.expectedBaseUrl = expectedBaseUrl; + this.expectedAuthKey = expectedAuthKey; + } + + @Override + public ProviderLoginResult register(String channelName, String baseUrl, String authKey) { + acceptedChannelName = expectedChannelName.equals(channelName); + acceptedBaseUrl = expectedBaseUrl.equals(baseUrl); + acceptedAuthKey = expectedAuthKey.equals(authKey); + registered.countDown(); + return new ProviderLoginResult(channelName, List.of(model(channelName, "alpha"))); + } + } + private static final class RecordingCore implements AgentCorePort { private final List requests = new ArrayList<>(); diff --git a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/LoginOverlayTest.java b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/LoginOverlayTest.java new file mode 100644 index 00000000..b02ec7f7 --- /dev/null +++ b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/LoginOverlayTest.java @@ -0,0 +1,81 @@ +package cn.lypi.transport.tui; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.List; +import org.junit.jupiter.api.Test; + +class LoginOverlayTest { + @Test + void collectsChannelThenBaseUrlAndMasksAuthKey() { + String authKey = "test-secret"; + LoginOverlay overlay = new LoginOverlay(); + + overlay.append("zen"); + assertTrue(overlay.accept().isEmpty()); + overlay.append("https://api.example.test/v1"); + assertTrue(overlay.accept().isEmpty()); + overlay.append(authKey); + + assertEquals(List.of( + "Channel name: zen", + "Base URL: https://api.example.test/v1", + "Auth key: ***********" + ), overlay.lines()); + assertFalse(String.join("\n", overlay.lines()).contains(authKey)); + + LoginOverlay.Submission submission = overlay.accept().orElseThrow(); + + assertEquals("zen", submission.channelName()); + assertEquals("https://api.example.test/v1", submission.baseUrl()); + assertTrue(authKey.equals(submission.authKey())); + assertFalse(submission.toString().contains(authKey)); + } + + @Test + void supportsBackspaceAndClearWithoutRetainingMaskedInput() { + LoginOverlay overlay = new LoginOverlay(); + overlay.append("zen"); + overlay.accept(); + overlay.append("https://api.example.test/v1"); + overlay.accept(); + overlay.append("secret"); + overlay.backspace(); + + assertEquals(List.of( + "Channel name: zen", + "Base URL: https://api.example.test/v1", + "Auth key: *****" + ), overlay.lines()); + + overlay.clear(); + + assertEquals(List.of(), overlay.lines()); + assertTrue(overlay.accept().isEmpty()); + } + + @Test + void refusesBlankInputAtEveryStep() { + LoginOverlay overlay = new LoginOverlay(); + overlay.open(); + + assertTrue(overlay.accept().isEmpty()); + assertEquals(List.of("Channel name: "), overlay.lines()); + + overlay.append("zen"); + assertTrue(overlay.accept().isEmpty()); + assertTrue(overlay.accept().isEmpty()); + assertEquals(List.of("Channel name: zen", "Base URL: "), overlay.lines()); + + overlay.append("https://api.example.test/v1"); + assertTrue(overlay.accept().isEmpty()); + assertTrue(overlay.accept().isEmpty()); + assertEquals(List.of( + "Channel name: zen", + "Base URL: https://api.example.test/v1", + "Auth key: " + ), overlay.lines()); + } +} diff --git a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/ModelPickerTest.java b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/ModelPickerTest.java new file mode 100644 index 00000000..091368d3 --- /dev/null +++ b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/ModelPickerTest.java @@ -0,0 +1,99 @@ +package cn.lypi.transport.tui; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import cn.lypi.contracts.model.ApiStyle; +import cn.lypi.contracts.model.CostProfile; +import cn.lypi.contracts.model.ModelDescriptor; +import cn.lypi.contracts.model.ModelSelection; +import cn.lypi.contracts.model.ThinkingLevel; +import java.math.BigDecimal; +import java.net.URI; +import java.util.List; +import java.util.Map; +import org.junit.jupiter.api.Test; + +class ModelPickerTest { + @Test + void sortsLabelsAndStartsAtCurrentModel() { + ModelPicker picker = new ModelPicker( + List.of(descriptor("zen", "kimi-k2.6"), descriptor("openai", "gpt-5-mini")), + new ModelSelection("openai", "gpt-5-mini", ThinkingLevel.MEDIUM) + ); + + assertEquals(List.of("openai/gpt-5-mini", "zen/kimi-k2.6"), picker.labels()); + assertEquals("openai/gpt-5-mini", picker.selectedLabel().orElseThrow()); + + picker.moveDown(); + + ModelDescriptor selected = picker.accept().orElseThrow(); + assertEquals("zen", selected.provider()); + assertEquals("kimi-k2.6", selected.modelId()); + } + + @Test + void deduplicatesProviderAndModelIdKeepingFirstDescriptor() { + ModelDescriptor first = descriptor("openai", "gpt-5-mini"); + ModelDescriptor duplicate = descriptor("openai", "gpt-5-mini"); + + ModelPicker picker = new ModelPicker(List.of(first, duplicate), null); + + assertEquals(List.of("openai/gpt-5-mini"), picker.labels()); + assertSame(first, picker.accept().orElseThrow()); + } + + @Test + void emptyCandidatesHaveNoSelection() { + ModelPicker picker = new ModelPicker(List.of(), null); + + picker.moveUp(); + picker.moveDown(); + + assertTrue(picker.labels().isEmpty()); + assertEquals(0, picker.selectedIndex()); + assertFalse(picker.selectedLabel().isPresent()); + assertFalse(picker.accept().isPresent()); + } + + @Test + void movingUpFromFirstWrapsToLast() { + ModelPicker picker = new ModelPicker( + List.of(descriptor("openai", "gpt-5-mini"), descriptor("zen", "kimi-k2.6")), + new ModelSelection("openai", "gpt-5-mini", ThinkingLevel.OFF) + ); + + picker.moveUp(); + + assertEquals("zen/kimi-k2.6", picker.selectedLabel().orElseThrow()); + } + + @Test + void movingDownFromLastWrapsToFirst() { + ModelPicker picker = new ModelPicker( + List.of(descriptor("openai", "gpt-5-mini"), descriptor("zen", "kimi-k2.6")), + new ModelSelection("zen", "kimi-k2.6", ThinkingLevel.HIGH) + ); + + picker.moveDown(); + + assertEquals("openai/gpt-5-mini", picker.selectedLabel().orElseThrow()); + } + + private static ModelDescriptor descriptor(String provider, String modelId) { + return new ModelDescriptor( + provider, + modelId, + URI.create("https://api.example.test/v1"), + ApiStyle.OPENAI_COMPATIBLE, + 128_000, + 16_384, + true, + false, + new CostProfile(BigDecimal.ZERO, BigDecimal.ZERO, "USD"), + Map.of() + ); + } +} diff --git a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/RuntimeTuiSubmitHandlerTest.java b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/RuntimeTuiSubmitHandlerTest.java index 602cc300..51954c91 100644 --- a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/RuntimeTuiSubmitHandlerTest.java +++ b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/RuntimeTuiSubmitHandlerTest.java @@ -24,6 +24,11 @@ import cn.lypi.contracts.event.MessageStartEvent; import cn.lypi.contracts.event.PermissionResponseEvent; import cn.lypi.contracts.event.SessionStateEvent; +import cn.lypi.contracts.error.ErrorSeverity; +import cn.lypi.contracts.error.ModelProviderException; +import cn.lypi.contracts.model.ApiStyle; +import cn.lypi.contracts.model.CostProfile; +import cn.lypi.contracts.model.ModelDescriptor; import cn.lypi.contracts.model.ModelSelection; import cn.lypi.contracts.model.ThinkingLevel; import cn.lypi.contracts.prompt.PromptParameter; @@ -35,6 +40,8 @@ import cn.lypi.contracts.runtime.AgentCorePort; import cn.lypi.contracts.runtime.CompactionRequest; import cn.lypi.contracts.runtime.CompactionResult; +import cn.lypi.contracts.runtime.ProviderLoginPort; +import cn.lypi.contracts.runtime.ProviderLoginResult; import cn.lypi.contracts.runtime.ResourceRuntimePort; import cn.lypi.contracts.runtime.SessionManagerPort; import cn.lypi.contracts.security.AgentMode; @@ -56,6 +63,7 @@ import cn.lypi.contracts.tui.SlashCommand; import cn.lypi.contracts.tui.SlashCommandHandler; import java.math.BigDecimal; +import java.net.URI; import java.nio.file.Path; import java.util.ArrayList; import java.util.List; @@ -451,6 +459,142 @@ void submitPermissionOptionPublishesResponseEvent() { assertEquals(false, event.fromKeyboardCancel()); } + @Test + void providerLoginRunsAsynchronouslyWithoutCreatingTurnOrSessionEntry() { + String authKey = "test-secret"; + RecordingCore core = new RecordingCore(); + RecordingEventBus events = new RecordingEventBus(); + RecordingSessionManager session = new RecordingSessionManager(); + QueuedExecutor executor = new QueuedExecutor(); + AtomicBoolean acceptedChannelName = new AtomicBoolean(); + AtomicBoolean acceptedBaseUrl = new AtomicBoolean(); + AtomicBoolean acceptedAuthKey = new AtomicBoolean(); + ProviderLoginPort login = (channelName, baseUrl, key) -> { + acceptedChannelName.set("zen".equals(channelName)); + acceptedBaseUrl.set("https://api.example.test/v1".equals(baseUrl)); + acceptedAuthKey.set(authKey.equals(key)); + return loginResult(); + }; + RuntimeTuiSubmitHandler handler = new RuntimeTuiSubmitHandler( + "ses_1", + core, + events, + executor, + new SlashCommandRouter("ses_1", Path.of("."), session, emptyResources()), + null, + skills(), + login + ); + + handler.submitProviderLogin("zen", "https://api.example.test/v1", authKey); + + assertEquals(1, executor.size()); + assertTrue(core.requests.isEmpty()); + assertTrue(session.entries.isEmpty()); + + executor.runNext(); + + assertTrue(acceptedChannelName.get()); + assertTrue(acceptedBaseUrl.get()); + assertTrue(acceptedAuthKey.get()); + assertEquals("login: registered login-example (1 model)", systemMessages(events).getFirst()); + assertFalse(events.published.stream().anyMatch(event -> event.toString().contains(authKey))); + assertTrue(core.requests.isEmpty()); + assertTrue(session.entries.isEmpty()); + } + + @Test + void providerLoginRejectsConcurrentSubmissionAndRedactsFailureMessages() { + String authKey = "test-secret"; + RecordingCore core = new RecordingCore(); + RecordingEventBus events = new RecordingEventBus(); + QueuedExecutor executor = new QueuedExecutor(); + AtomicInteger registrations = new AtomicInteger(); + ProviderLoginPort login = (channelName, baseUrl, key) -> { + registrations.incrementAndGet(); + throw new ModelProviderException( + "provider.login_failed", + ErrorSeverity.ERROR, + false, + "provider rejected " + key + ); + }; + RuntimeTuiSubmitHandler handler = new RuntimeTuiSubmitHandler( + "ses_1", + core, + events, + executor, + null, + null, + skills(), + login + ); + + handler.submitProviderLogin("zen", "https://api.example.test/v1", authKey); + handler.submitProviderLogin("zen", "https://api.example.test/v1", authKey); + + assertEquals(1, executor.size()); + ErrorEvent concurrent = assertInstanceOf(ErrorEvent.class, events.published.getFirst()); + assertEquals("login: provider registration is running", concurrent.message()); + + executor.runNext(); + + assertEquals(1, registrations.get()); + ErrorEvent failure = assertInstanceOf(ErrorEvent.class, events.published.getLast()); + assertEquals("login: provider registration failed", failure.message()); + assertFalse(failure.message().contains(authKey)); + assertTrue(core.requests.isEmpty()); + } + + @Test + void unavailableProviderLoginProducesFixedErrorWithoutStartingTurn() { + String authKey = "test-secret"; + RecordingCore core = new RecordingCore(); + RecordingEventBus events = new RecordingEventBus(); + QueuedExecutor executor = new QueuedExecutor(); + RuntimeTuiSubmitHandler handler = new RuntimeTuiSubmitHandler("ses_1", core, events, executor); + + handler.submitProviderLogin("zen", "https://api.example.test/v1", authKey); + executor.runNext(); + + ErrorEvent error = assertInstanceOf(ErrorEvent.class, events.published.getFirst()); + assertEquals("login: provider login is unavailable", error.message()); + assertFalse(error.message().contains(authKey)); + assertTrue(core.requests.isEmpty()); + } + + @Test + void providerLoginShowsSanitizedProviderError() { + RecordingCore core = new RecordingCore(); + RecordingEventBus events = new RecordingEventBus(); + QueuedExecutor executor = new QueuedExecutor(); + ProviderLoginPort login = (channelName, baseUrl, authKey) -> { + throw new ModelProviderException( + "provider.login_invalid_auth_key", + ErrorSeverity.ERROR, + false, + "Provider auth key is required." + ); + }; + RuntimeTuiSubmitHandler handler = new RuntimeTuiSubmitHandler( + "ses_1", + core, + events, + executor, + null, + null, + skills(), + login + ); + + handler.submitProviderLogin("zen", "https://api.example.test/v1", "test-secret"); + executor.runNext(); + + ErrorEvent error = assertInstanceOf(ErrorEvent.class, events.published.getFirst()); + assertEquals("login: Provider auth key is required.", error.message()); + assertTrue(core.requests.isEmpty()); + } + @Test void stateSlashCommandDoesNotSubmitTurnAndAppendsSessionEntry() { RecordingCore core = new RecordingCore(); @@ -1040,6 +1184,29 @@ private static ResourceRuntimePort emptyResources() { return resources(List.of()); } + private static ProviderLoginResult loginResult() { + return new ProviderLoginResult("login-example", List.of(new ModelDescriptor( + "login-example", + "alpha", + URI.create("https://api.example.test/v1"), + ApiStyle.OPENAI_COMPATIBLE, + 0, + 0, + false, + false, + new CostProfile(BigDecimal.ZERO, BigDecimal.ZERO, "USD"), + Map.of() + ))); + } + + private static List systemMessages(RecordingEventBus events) { + return events.published.stream() + .filter(MessageDeltaEvent.class::isInstance) + .map(MessageDeltaEvent.class::cast) + .map(MessageDeltaEvent::delta) + .toList(); + } + private static ResourceRuntimePort reviewResources() { PromptTemplate review = new PromptTemplate( "review", diff --git a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/SlashCommandPickerTest.java b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/SlashCommandPickerTest.java index 8769a45c..092dccdf 100644 --- a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/SlashCommandPickerTest.java +++ b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/SlashCommandPickerTest.java @@ -14,7 +14,7 @@ void defaultCommandsIncludeImplementedStateCommandsAndTemplates() { picker.updateFilter("/"); assertEquals( - List.of("/model", "/thinking", "/plan", "/permission-mode", "/compact", "/review", "/commit"), + List.of("/model", "/login", "/thinking", "/plan", "/permission-mode", "/compact", "/review", "/commit"), picker.visibleCommands() ); } diff --git a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/SlashCommandRouterTest.java b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/SlashCommandRouterTest.java index a8955d7e..64568a10 100644 --- a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/SlashCommandRouterTest.java +++ b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/SlashCommandRouterTest.java @@ -7,6 +7,10 @@ import cn.lypi.contracts.context.AgentMessage; import cn.lypi.contracts.context.ContextBudget; +import cn.lypi.contracts.model.ApiStyle; +import cn.lypi.contracts.model.CostProfile; +import cn.lypi.contracts.model.ModelCatalogPort; +import cn.lypi.contracts.model.ModelDescriptor; import cn.lypi.contracts.model.ModelSelection; import cn.lypi.contracts.model.ThinkingLevel; import cn.lypi.contracts.prompt.PromptParameter; @@ -34,6 +38,7 @@ import cn.lypi.contracts.tui.NewSessionController; import cn.lypi.contracts.tui.SessionRuntimeState; import java.math.BigDecimal; +import java.net.URI; import java.nio.file.Path; import java.util.ArrayList; import java.util.LinkedHashMap; @@ -128,6 +133,85 @@ void singleModelArgumentKeepsCurrentProvider() { assertEquals(new ModelSelection("openai", "gpt-5.4", ThinkingLevel.HIGH), model.model()); } + @Test + void catalogAllowsKnownQualifiedModelAndPreservesThinkingLevel() { + RecordingSessionManager session = new RecordingSessionManager(context( + new ModelSelection("openai", "gpt-5", ThinkingLevel.MEDIUM), + ThinkingLevel.LOW, + AgentMode.EXECUTE, + PermissionMode.ASK + )); + SlashCommandRouter router = new SlashCommandRouter( + "ses_1", + Path.of("."), + session, + emptyResources(), + null, + null, + List.of(), + catalog(model("zen", "kimi-k2.6")) + ); + + SlashCommandResult result = router.route("/model zen/kimi-k2.6"); + + ModelChangeEntry entry = assertInstanceOf(ModelChangeEntry.class, session.entries.getFirst()); + assertEquals(new ModelSelection("zen", "kimi-k2.6", ThinkingLevel.LOW), entry.model()); + assertEquals("model: zen/kimi-k2.6", result.notice().orElseThrow()); + } + + @Test + void catalogSelectionTurnsThinkingOffForModelsThatDoNotSupportIt() { + RecordingSessionManager session = new RecordingSessionManager(context( + new ModelSelection("openai", "gpt-5", ThinkingLevel.MEDIUM), + ThinkingLevel.HIGH, + AgentMode.EXECUTE, + PermissionMode.ASK + )); + SlashCommandRouter router = new SlashCommandRouter( + "ses_1", + Path.of("."), + session, + emptyResources(), + null, + null, + List.of(), + catalog(model("login-fixture", "discovered-model", false)) + ); + + router.route("/model login-fixture/discovered-model"); + + ModelChangeEntry entry = assertInstanceOf(ModelChangeEntry.class, session.entries.getFirst()); + assertEquals( + new ModelSelection("login-fixture", "discovered-model", ThinkingLevel.OFF), + entry.model() + ); + } + + @Test + void catalogRejectsUnknownModelWithoutAppendingEntry() { + RecordingSessionManager session = new RecordingSessionManager(context( + new ModelSelection("openai", "gpt-5", ThinkingLevel.MEDIUM), + ThinkingLevel.MEDIUM, + AgentMode.EXECUTE, + PermissionMode.ASK + )); + SlashCommandRouter router = new SlashCommandRouter( + "ses_1", + Path.of("."), + session, + emptyResources(), + null, + null, + List.of(), + catalog(model("zen", "kimi-k2.6")) + ); + + SlashCommandResult result = router.route("/model zen/not-listed"); + + assertEquals("unknown model: zen/not-listed", result.message().orElseThrow()); + assertEquals(List.of(), session.entries); + } + @Test void invalidModelProviderSyntaxIsConsumedWithErrorButDoesNotAppend() { RecordingSessionManager session = new RecordingSessionManager(context( @@ -251,6 +335,26 @@ void unknownSlashCommandFallsThroughToModel() { assertEquals(List.of(), session.entries); } + @Test + void loginCommandIsBuiltInAndRejectsArgumentsWithoutChangingSessionState() { + RecordingSessionManager session = new RecordingSessionManager(context( + new ModelSelection("openai", "gpt-5", ThinkingLevel.MEDIUM), + ThinkingLevel.MEDIUM, + AgentMode.EXECUTE, + PermissionMode.ASK + )); + SlashCommandRouter router = new SlashCommandRouter("ses_1", Path.of("."), session, emptyResources()); + + SlashCommandResult exact = router.route("/login"); + SlashCommandResult withArgument = router.route("/login https://example.test/v1"); + + assertTrue(router.commandNames().contains("/login")); + assertTrue(exact.matched()); + assertTrue(exact.consumed()); + assertEquals("usage: /login", withArgument.message().orElseThrow()); + assertEquals(List.of(), session.entries); + } + @Test void removedModeCommandDoesNotPrefixMatchModelCommand() { RecordingSessionManager session = new RecordingSessionManager(context( @@ -548,6 +652,43 @@ private static SessionContext context( return new SessionContext(List.of(), List.of("root"), List.of(), model, thinking, mode, permissionMode); } + private static ModelCatalogPort catalog(ModelDescriptor... descriptors) { + List models = List.of(descriptors); + return new ModelCatalogPort() { + @Override + public List list() { + return models; + } + + @Override + public Optional find(ModelSelection selection) { + return models.stream() + .filter(candidate -> candidate.provider().equals(selection.provider())) + .filter(candidate -> candidate.modelId().equals(selection.modelId())) + .findFirst(); + } + }; + } + + private static ModelDescriptor model(String provider, String modelId) { + return model(provider, modelId, true); + } + + private static ModelDescriptor model(String provider, String modelId, boolean supportsThinking) { + return new ModelDescriptor( + provider, + modelId, + URI.create("https://api.example.test/v1"), + ApiStyle.OPENAI_COMPATIBLE, + 128_000, + 16_384, + supportsThinking, + false, + new CostProfile(BigDecimal.ZERO, BigDecimal.ZERO, "USD"), + Map.of() + ); + } + private static ResourceRuntimePort emptyResources() { return resourcesWith(); } diff --git a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/TuiInputLoopTest.java b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/TuiInputLoopTest.java index 9990c74b..c1e2cc30 100644 --- a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/TuiInputLoopTest.java +++ b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/TuiInputLoopTest.java @@ -12,6 +12,11 @@ import cn.lypi.contracts.event.MessageBlockSnapshot; import cn.lypi.contracts.event.MessageEndEvent; import cn.lypi.contracts.event.MessageStartEvent; +import cn.lypi.contracts.model.ApiStyle; +import cn.lypi.contracts.model.CostProfile; +import cn.lypi.contracts.model.ModelDescriptor; +import cn.lypi.contracts.model.ModelSelection; +import cn.lypi.contracts.model.ThinkingLevel; import cn.lypi.contracts.tui.BranchSummaryOffer; import cn.lypi.contracts.tui.PermissionPromptView; import cn.lypi.contracts.tui.ResumeSessionController; @@ -34,14 +39,17 @@ import cn.lypi.contracts.skill.SkillIndex; import cn.lypi.contracts.skill.SkillMention; import cn.lypi.contracts.skill.SkillSource; +import java.math.BigDecimal; +import java.net.URI; +import java.nio.file.Path; +import java.time.Instant; import java.util.ArrayList; import java.util.List; import java.util.Map; import java.util.Optional; +import java.util.concurrent.atomic.AtomicReference; import java.util.function.Consumer; import java.util.function.Supplier; -import java.nio.file.Path; -import java.time.Instant; import org.junit.jupiter.api.Test; class TuiInputLoopTest { @@ -125,6 +133,30 @@ private static TuiInputLoop testLoop( ResumeSessionController resumeController, Consumer resumeStateConsumer, Supplier skillIndexSupplier + ) { + return testLoop( + submitHandler, + frameConsumer, + layout, + viewSupplier, + slashPickerSupplier, + resumeController, + resumeStateConsumer, + skillIndexSupplier, + null + ); + } + + private static TuiInputLoop testLoop( + TuiSubmitHandler submitHandler, + Consumer> frameConsumer, + TuiLayout layout, + Supplier viewSupplier, + Supplier slashPickerSupplier, + ResumeSessionController resumeController, + Consumer resumeStateConsumer, + Supplier skillIndexSupplier, + Supplier modelPickerSupplier ) { TestRenderRequest renderRequest = new TestRenderRequest(frameConsumer, layout); TuiInputLoop loop = new TuiInputLoop( @@ -135,12 +167,33 @@ private static TuiInputLoop testLoop( slashPickerSupplier, resumeController, resumeStateConsumer, - skillIndexSupplier + skillIndexSupplier, + modelPickerSupplier ); renderRequest.bind(loop); return loop; } + private static TuiInputLoop modelLoop( + TuiSubmitHandler submitHandler, + Consumer> frameConsumer, + TuiLayout layout, + Supplier viewSupplier, + Supplier modelPickerSupplier + ) { + return testLoop( + submitHandler, + frameConsumer, + layout, + viewSupplier, + () -> new SlashCommandPicker(List.of("/model")), + null, + null, + null, + modelPickerSupplier + ); + } + private static final class TestRenderRequest implements Runnable { private final Consumer> frameConsumer; private final TuiLayout layout; @@ -822,7 +875,7 @@ void slashOverlayShowsCandidatesAndAcceptsSelection() { loop.acceptText("/"); assertTrue(frames.getLast().contains("> /model")); - assertTrue(frames.getLast().contains(" /compact")); + assertTrue(frames.getLast().contains(" /login")); loop.acceptText("th"); assertTrue(frames.getLast().contains("> /thinking")); @@ -833,6 +886,249 @@ void slashOverlayShowsCandidatesAndAcceptsSelection() { assertEquals(List.of(), submit.submitted); } + @Test + void loginOverlayMasksAuthKeyAndSubmitsOutsideEditorHistory() { + String authKey = "test-secret"; + RecordingSubmitHandler submit = new RecordingSubmitHandler(); + TuiInputLoop loop = testLoop( + submit, + ignored -> { + }, + new TuiLayout(40, 9) + ); + + loop.acceptText("ordinary input"); + loop.acceptKey(TerminalKey.ENTER); + loop.acceptText("/login"); + loop.acceptKey(TerminalKey.ENTER); + loop.acceptPaste("zen"); + loop.acceptKey(TerminalKey.ENTER); + loop.acceptPaste("https://api.example.test/v1"); + loop.acceptKey(TerminalKey.ENTER); + loop.acceptPaste(authKey); + + assertEquals("", loop.draft()); + assertEquals(List.of( + "Channel name: zen", + "Base URL: https://api.example.test/v1", + "Auth key: ***********" + ), loop.overlayLines()); + assertFalse(String.join("\n", loop.overlayLines()).contains(authKey)); + + loop.acceptKey(TerminalKey.ENTER); + + assertEquals(List.of("ordinary input"), submit.submitted); + assertEquals(1, submit.providerLogins.size()); + assertEquals("zen", submit.providerLogins.getFirst().channelName()); + assertEquals("https://api.example.test/v1", submit.providerLogins.getFirst().baseUrl()); + assertTrue(authKey.equals(submit.providerLogins.getFirst().authKey())); + assertEquals(List.of(), loop.overlayLines()); + loop.acceptKey(TerminalKey.UP); + assertEquals("ordinary input", loop.draft()); + } + + @Test + void escapeAndCtrlCCancelLoginOverlayWithoutInterruptingTheActiveTurn() { + RecordingSubmitHandler submit = new RecordingSubmitHandler(); + TuiInputLoop loop = testLoop( + submit, + ignored -> { + }, + new TuiLayout(40, 9) + ); + + loop.acceptText("/login"); + loop.acceptKey(TerminalKey.ENTER); + loop.acceptText("zen"); + loop.acceptKey(TerminalKey.ESC); + loop.acceptText("/login"); + loop.acceptKey(TerminalKey.ENTER); + loop.acceptText("zen"); + loop.acceptKey(TerminalKey.CTRL_C); + + assertEquals(List.of(), loop.overlayLines()); + assertEquals(List.of(), submit.providerLogins); + assertEquals(0, submit.interrupts); + assertEquals(0, submit.exits); + } + + @Test + void modelSlashOpensPickerAndSubmitsSelectedModel() { + RecordingSubmitHandler submit = new RecordingSubmitHandler(); + TuiInputLoop loop = modelLoop( + submit, + ignored -> { + }, + new TuiLayout(40, 9), + null, + () -> new ModelPicker( + List.of(model("zen", "kimi-k2.6"), model("openai", "gpt-5-mini")), + new ModelSelection("openai", "gpt-5-mini", ThinkingLevel.MEDIUM) + ) + ); + + loop.acceptText("/model"); + loop.acceptKey(TerminalKey.ENTER); + + assertEquals(List.of("> openai/gpt-5-mini", " zen/kimi-k2.6"), loop.overlayLines()); + assertEquals(List.of(), submit.submitted); + + loop.acceptKey(TerminalKey.DOWN); + loop.acceptKey(TerminalKey.ENTER); + + assertEquals(List.of("/model zen/kimi-k2.6"), submit.submitted); + assertEquals(List.of(), loop.overlayLines()); + assertEquals("", loop.draft()); + } + + @Test + void exactModelDraftOpensPickerAfterSlashOverlayWasClosed() { + RecordingSubmitHandler submit = new RecordingSubmitHandler(); + TuiInputLoop loop = modelLoop( + submit, + ignored -> { + }, + new TuiLayout(40, 9), + null, + () -> new ModelPicker(List.of(model("openai", "gpt-5-mini")), null) + ); + + loop.acceptText("/model"); + loop.acceptKey(TerminalKey.ESC); + loop.acceptKey(TerminalKey.ENTER); + + assertEquals(List.of("> openai/gpt-5-mini"), loop.overlayLines()); + assertEquals(List.of(), submit.submitted); + assertEquals("", loop.draft()); + } + + @Test + void escapeClosesModelPickerWithoutSubmitting() { + RecordingSubmitHandler submit = new RecordingSubmitHandler(); + TuiInputLoop loop = modelLoop( + submit, + ignored -> { + }, + new TuiLayout(40, 9), + null, + () -> new ModelPicker(List.of(model("openai", "gpt-5-mini")), null) + ); + + loop.acceptText("/model"); + loop.acceptKey(TerminalKey.ENTER); + loop.acceptKey(TerminalKey.ESC); + + assertEquals(List.of(), loop.overlayLines()); + assertEquals(List.of(), submit.submitted); + assertEquals("", loop.draft()); + } + + @Test + void emptyModelPickerStaysOpenOnEnter() { + RecordingSubmitHandler submit = new RecordingSubmitHandler(); + TuiInputLoop loop = modelLoop( + submit, + ignored -> { + }, + new TuiLayout(40, 9), + null, + () -> new ModelPicker(List.of(), null) + ); + + loop.acceptText("/model"); + loop.acceptKey(TerminalKey.ENTER); + + assertEquals(List.of("No models available"), loop.overlayLines()); + + loop.acceptKey(TerminalKey.ENTER); + + assertEquals(List.of("No models available"), loop.overlayLines()); + assertEquals(List.of(), submit.submitted); + } + + @Test + void modelPickerBlocksTextAndPasteUntilClosed() { + RecordingSubmitHandler submit = new RecordingSubmitHandler(); + TuiInputLoop loop = modelLoop( + submit, + ignored -> { + }, + new TuiLayout(40, 9), + null, + () -> new ModelPicker(List.of(model("openai", "gpt-5-mini")), null) + ); + + loop.acceptText("/model"); + loop.acceptKey(TerminalKey.ENTER); + loop.acceptText("ignored"); + loop.acceptPaste("also ignored"); + + assertEquals("", loop.draft()); + assertEquals(List.of("> openai/gpt-5-mini"), loop.overlayLines()); + } + + @Test + void permissionPromptTakesPriorityOverOpenModelPicker() { + RecordingSubmitHandler submit = new RecordingSubmitHandler(); + AtomicReference view = new AtomicReference<>(null); + TuiInputLoop loop = modelLoop( + submit, + ignored -> { + }, + new TuiLayout(40, 9), + () -> view.get() == null + ? new TuiViewModel( + List.of(), + new StatusBarState("ses_1", "gpt-5-mini", "ready", "ASK"), + List.of(), + Optional.empty(), + Optional.empty() + ) + : view.get(), + () -> new ModelPicker( + List.of(model("openai", "gpt-5-mini"), model("zen", "kimi-k2.6")), + null + ) + ); + + loop.acceptText("/model"); + loop.acceptKey(TerminalKey.ENTER); + view.set(permissionViewWithOptions("allow_once", "escape_cancel")); + + loop.acceptKey(TerminalKey.DOWN); + loop.acceptKey(TerminalKey.ENTER); + + assertEquals(List.of("perm_toolu_1:toolu_1:remember"), submit.permissionOptions); + assertEquals(List.of(), submit.submitted); + assertEquals(List.of(), loop.overlayLines()); + } + + @Test + void modelPickerScrollsSelectedModelIntoVisibleWindow() { + RecordingSubmitHandler submit = new RecordingSubmitHandler(); + List models = java.util.stream.IntStream.range(0, 8) + .mapToObj(index -> model("provider", "model-%02d".formatted(index))) + .toList(); + TuiInputLoop loop = modelLoop( + submit, + ignored -> { + }, + new TuiLayout(40, 7), + null, + () -> new ModelPicker(models, null) + ); + + loop.acceptText("/model"); + loop.acceptKey(TerminalKey.ENTER); + for (int index = 0; index < 5; index++) { + loop.acceptKey(TerminalKey.DOWN); + } + + assertEquals(3, loop.overlayLines().size()); + assertTrue(loop.overlayLines().contains("> provider/model-05")); + assertFalse(loop.overlayLines().stream().anyMatch(line -> line.contains("provider/model-00"))); + } + @Test void slashOverlayUsesArrowKeysAndEscWithoutHistoryNavigation() { RecordingSubmitHandler submit = new RecordingSubmitHandler(); @@ -1382,6 +1678,21 @@ private static String inputContent(String content) { return INPUT_BACKGROUND + content + ANSI_RESET; } + private static ModelDescriptor model(String provider, String modelId) { + return new ModelDescriptor( + provider, + modelId, + URI.create("https://api.example.test/v1"), + ApiStyle.OPENAI_COMPATIBLE, + 128_000, + 16_384, + true, + false, + new CostProfile(BigDecimal.ZERO, BigDecimal.ZERO, "USD"), + Map.of() + ); + } + private static ResumeSessionController emptyResumeController() { return new ResumeSessionController() { @Override @@ -1460,6 +1771,7 @@ private static final class RecordingSubmitHandler implements TuiSubmitHandler { private final List permissionOptions = new ArrayList<>(); private final List resumes = new ArrayList<>(); private final List interruptReasons = new ArrayList<>(); + private final List providerLogins = new ArrayList<>(); private int interrupts; private int exits; @@ -1475,6 +1787,11 @@ public void submitUserInput(String input, List skillMentions) { this.skillMentions.add(skillMentions); } + @Override + public void submitProviderLogin(String channelName, String baseUrl, String authKey) { + providerLogins.add(new LoginSubmission(channelName, baseUrl, authKey)); + } + @Override public List pendingSteeringMessages() { return List.copyOf(pendingSteering); @@ -1516,6 +1833,14 @@ public void submitPermissionOption(String requestId, String toolUseId, String op public void resumeSession(String sessionId, String leafId) { resumes.add(sessionId + ":" + leafId); } + + private record LoginSubmission(String channelName, String baseUrl, String authKey) { + @Override + public String toString() { + return "LoginSubmission[channelName=" + channelName + + ", baseUrl=" + baseUrl + ", authKey=]"; + } + } } private static SkillIndex skills(String name, String description) {