From 8381c3c2dadd0e7d461a0ce4474b21c084a19b64 Mon Sep 17 00:00:00 2001 From: liujihui Date: Tue, 4 Aug 2026 11:37:11 +0800 Subject: [PATCH 1/2] feat: add Java client module and CI --- .github/workflows/citest.yaml | 5 + .github/workflows/citest_npu.yaml | 5 + .github/workflows/java-ci.yaml | 36 +++ .github/workflows/lint.yaml | 10 +- javaclients/.editorconfig | 19 ++ javaclients/.gitignore | 6 + javaclients/CHANGELOG.md | 10 + javaclients/CONTRIBUTING.md | 7 + javaclients/LICENSE | 63 +++++ javaclients/README.md | 82 ++++++ javaclients/SECURITY.md | 5 + javaclients/pom.xml | 65 +++++ .../modelscope/twinkle/TwinkleClient.java | 151 ++++++++++ .../twinkle/config/ClientConfig.java | 117 ++++++++ .../twinkle/exception/TwinkleException.java | 13 + .../TwinkleIterationExhaustedException.java | 11 + .../exception/TwinkleProtocolException.java | 9 + .../exception/TwinkleServiceException.java | 36 +++ .../exception/TwinkleTransportException.java | 9 + .../twinkle/internal/ResponseMapper.java | 104 +++++++ .../modelscope/twinkle/model/ModelClient.java | 258 ++++++++++++++++++ .../twinkle/model/ModelsClient.java | 17 ++ .../twinkle/processor/DataLoaderClient.java | 67 +++++ .../twinkle/processor/DatasetClient.java | 138 ++++++++++ .../processor/InputProcessorClient.java | 24 ++ .../twinkle/processor/ProcessorsClient.java | 53 ++++ .../twinkle/processor/RemoteProcessor.java | 34 +++ .../twinkle/runs/TrainingRunsClient.java | 160 +++++++++++ .../twinkle/sampler/SamplerClient.java | 124 +++++++++ .../twinkle/sampler/SamplersClient.java | 17 ++ .../twinkle/session/SessionManager.java | 68 +++++ .../twinkle/transport/HttpTransport.java | 16 ++ .../twinkle/transport/JsonValue.java | 11 + .../twinkle/transport/OkHttpTransport.java | 147 ++++++++++ .../twinkle/transport/TwinkleJsonCodec.java | 50 ++++ .../transport/TwinkleSerializable.java | 8 + .../twinkle/types/CapacityInfo.java | 12 + .../modelscope/twinkle/types/Checkpoint.java | 11 + .../twinkle/types/CheckpointPage.java | 12 + .../twinkle/types/CheckpointPath.java | 13 + .../modelscope/twinkle/types/Cursor.java | 9 + .../modelscope/twinkle/types/DatasetKind.java | 20 ++ .../modelscope/twinkle/types/DatasetMeta.java | 74 +++++ .../twinkle/types/DeleteCheckpointResult.java | 11 + .../modelscope/twinkle/types/LoraConfig.java | 41 +++ .../twinkle/types/SampleRequest.java | 21 ++ .../twinkle/types/SampleResult.java | 9 + .../twinkle/types/SampledSequence.java | 12 + .../twinkle/types/SaveResponse.java | 8 + .../twinkle/types/ServerCapabilities.java | 11 + .../twinkle/types/SupportedModel.java | 10 + .../modelscope/twinkle/types/TrainingRun.java | 11 + .../twinkle/types/TrainingRunPage.java | 12 + .../modelscope/twinkle/types/WeightsInfo.java | 13 + .../io/github/modelscope/twinkle/LjhTest.java | 116 ++++++++ .../twinkle/ProtocolContractTest.java | 62 +++++ .../modelscope/twinkle/TwinkleClientTest.java | 63 +++++ 57 files changed, 2505 insertions(+), 1 deletion(-) create mode 100644 .github/workflows/java-ci.yaml create mode 100644 javaclients/.editorconfig create mode 100644 javaclients/.gitignore create mode 100644 javaclients/CHANGELOG.md create mode 100644 javaclients/CONTRIBUTING.md create mode 100644 javaclients/LICENSE create mode 100644 javaclients/README.md create mode 100644 javaclients/SECURITY.md create mode 100644 javaclients/pom.xml create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/TwinkleClient.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/config/ClientConfig.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleException.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleIterationExhaustedException.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleProtocolException.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleServiceException.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleTransportException.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/internal/ResponseMapper.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/model/ModelClient.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/model/ModelsClient.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/processor/DataLoaderClient.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/processor/DatasetClient.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/processor/InputProcessorClient.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/processor/ProcessorsClient.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/processor/RemoteProcessor.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/runs/TrainingRunsClient.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/sampler/SamplerClient.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/sampler/SamplersClient.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/session/SessionManager.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/transport/HttpTransport.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/transport/JsonValue.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/transport/OkHttpTransport.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/transport/TwinkleJsonCodec.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/transport/TwinkleSerializable.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/CapacityInfo.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/Checkpoint.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/CheckpointPage.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/CheckpointPath.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/Cursor.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/DatasetKind.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/DatasetMeta.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/DeleteCheckpointResult.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/LoraConfig.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/SampleRequest.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/SampleResult.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/SampledSequence.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/SaveResponse.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/ServerCapabilities.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/SupportedModel.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/TrainingRun.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/TrainingRunPage.java create mode 100644 javaclients/src/main/java/io/github/modelscope/twinkle/types/WeightsInfo.java create mode 100644 javaclients/src/test/java/io/github/modelscope/twinkle/LjhTest.java create mode 100644 javaclients/src/test/java/io/github/modelscope/twinkle/ProtocolContractTest.java create mode 100644 javaclients/src/test/java/io/github/modelscope/twinkle/TwinkleClientTest.java diff --git a/.github/workflows/citest.yaml b/.github/workflows/citest.yaml index 4e16ebdc..1d217930 100644 --- a/.github/workflows/citest.yaml +++ b/.github/workflows/citest.yaml @@ -3,6 +3,7 @@ name: citest on: push: branches: + - main - master - "release/**" paths-ignore: @@ -16,6 +17,8 @@ on: - "NOTICE" - ".github/workflows/lint.yaml" - ".github/workflows/publish.yaml" + - "javaclients/**" + - ".github/workflows/java-ci.yaml" pull_request: paths-ignore: @@ -29,6 +32,8 @@ on: - "NOTICE" - ".github/workflows/lint.yaml" - ".github/workflows/publish.yaml" + - "javaclients/**" + - ".github/workflows/java-ci.yaml" concurrency: group: ${{ github.workflow }}-${{ github.ref }} diff --git a/.github/workflows/citest_npu.yaml b/.github/workflows/citest_npu.yaml index a3878dae..c27e7a0d 100644 --- a/.github/workflows/citest_npu.yaml +++ b/.github/workflows/citest_npu.yaml @@ -3,6 +3,7 @@ name: citest-npu on: push: branches: + - main - master - "release/**" paths-ignore: @@ -17,6 +18,8 @@ on: - "NOTICE" - ".github/workflows/lint.yaml" - ".github/workflows/publish.yaml" + - "javaclients/**" + - ".github/workflows/java-ci.yaml" pull_request: paths-ignore: @@ -31,6 +34,8 @@ on: - "NOTICE" - ".github/workflows/lint.yaml" - ".github/workflows/publish.yaml" + - "javaclients/**" + - ".github/workflows/java-ci.yaml" concurrency: group: ${{ github.workflow }}-${{ github.ref }} diff --git a/.github/workflows/java-ci.yaml b/.github/workflows/java-ci.yaml new file mode 100644 index 00000000..2b5fac44 --- /dev/null +++ b/.github/workflows/java-ci.yaml @@ -0,0 +1,36 @@ +name: Java Client CI + +on: + push: + paths: + - "javaclients/**" + - ".github/workflows/java-ci.yaml" + pull_request: + paths: + - "javaclients/**" + - ".github/workflows/java-ci.yaml" + workflow_dispatch: + +concurrency: + group: ${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +jobs: + verify: + name: Maven 构建与测试 + runs-on: ubuntu-latest + + steps: + - name: 拉取代码 + uses: actions/checkout@v4 + + - name: 配置 JDK 17 + uses: actions/setup-java@v4 + with: + distribution: temurin + java-version: "17" + cache: maven + cache-dependency-path: javaclients/pom.xml + + - name: Maven 构建与测试 + run: mvn -B -f javaclients/pom.xml verify diff --git a/.github/workflows/lint.yaml b/.github/workflows/lint.yaml index 771ee4bc..57ba5b9a 100644 --- a/.github/workflows/lint.yaml +++ b/.github/workflows/lint.yaml @@ -1,6 +1,14 @@ name: Lint test -on: [push, pull_request] +on: + push: + paths-ignore: + - "javaclients/**" + - ".github/workflows/java-ci.yaml" + pull_request: + paths-ignore: + - "javaclients/**" + - ".github/workflows/java-ci.yaml" concurrency: group: ${{ github.workflow }}-${{ github.ref }} diff --git a/javaclients/.editorconfig b/javaclients/.editorconfig new file mode 100644 index 00000000..1b8582ef --- /dev/null +++ b/javaclients/.editorconfig @@ -0,0 +1,19 @@ +root = true + +[*] +charset = utf-8 +end_of_line = lf +insert_final_newline = true +trim_trailing_whitespace = true + +[*.java] +indent_style = space +indent_size = 4 +ij_java_right_margin = 100 +ij_java_use_single_class_import = true +ij_java_class_count_to_use_import_on_demand = 999 +ij_java_names_count_to_use_import_on_demand = 999 + +[*.{xml,yml,yaml,md}] +indent_style = space +indent_size = 2 diff --git a/javaclients/.gitignore b/javaclients/.gitignore new file mode 100644 index 00000000..16e4f7cf --- /dev/null +++ b/javaclients/.gitignore @@ -0,0 +1,6 @@ +target/ +.idea/ +*.iml +.DS_Store +.env +*.log diff --git a/javaclients/CHANGELOG.md b/javaclients/CHANGELOG.md new file mode 100644 index 00000000..92160238 --- /dev/null +++ b/javaclients/CHANGELOG.md @@ -0,0 +1,10 @@ +# 更新日志 + +本项目遵循语义化版本规范。 + +## 1.0.0-SNAPSHOT + +- 新增 Java 17 的 Twinkle HTTP 客户端基础实现。 +- 新增会话、心跳、模型训练、采样、训练任务、数据集和远程处理器 API。 +- 与 Python `twinkle_client_simple` 对齐 `/twinkle` 管理路由、LoRA 和 DatasetMeta 的核心序列化协议。 +- 新增 MockWebServer 协议契约测试与环境变量驱动的服务联调示例。 diff --git a/javaclients/CONTRIBUTING.md b/javaclients/CONTRIBUTING.md new file mode 100644 index 00000000..d44b4408 --- /dev/null +++ b/javaclients/CONTRIBUTING.md @@ -0,0 +1,7 @@ +# 贡献指南 + +欢迎提交 Issue 和 Pull Request。提交前请使用 IntelliJ IDEA 或本地 Maven 完成编译检查,并保持公开文档、JavaDoc、示例说明和源代码注释为中文。 + +提交内容不得包含令牌、账号、内网地址、用户数据、本地绝对路径或 IDE 生成文件。API 标识符和协议字段使用英文,以维持 Java 和服务端协议兼容性。 + +建议的提交格式为:`feat:`、`fix:`、`docs:`、`build:` 或 `chore:`。 diff --git a/javaclients/LICENSE b/javaclients/LICENSE new file mode 100644 index 00000000..9cef33c6 --- /dev/null +++ b/javaclients/LICENSE @@ -0,0 +1,63 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of this + License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work. + + 4. Redistribution. You may reproduce and distribute copies of the Work + or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You meet + the following conditions: You must give any other recipients of the + Work or Derivative Works a copy of this License; cause modified files + to carry prominent notices; retain all copyright, patent, trademark, + and attribution notices; and include a readable copy of attribution + notices from any NOTICE file. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + shall be under the terms of this License, without additional terms. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor. + + 7. Disclaimer of Warranty. Unless required by applicable law or agreed + to in writing, Licensor provides the Work on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort, contract, or otherwise, shall any Contributor be + liable to You for damages arising as a result of this License or use + of the Work. + + 9. Accepting Warranty or Additional Liability. While redistributing the + Work, You may choose to offer support, warranty, indemnity, or other + liability obligations consistent with this License. diff --git a/javaclients/README.md b/javaclients/README.md new file mode 100644 index 00000000..f2cc2c77 --- /dev/null +++ b/javaclients/README.md @@ -0,0 +1,82 @@ +# Twinkle Java Client + +Twinkle 训练服务的 Java 17 同步客户端。该项目面向国内开发者发布:使用中文文档、中文 JavaDoc 与中文源代码注释;Java API 名称和 HTTP 协议字段保持英文,以便与 Java 生态及 Twinkle 服务端兼容。 + +## 特性 + +- 自动创建并维护服务端会话心跳,客户端关闭时自动停止。 +- 提供模型训练、LoRA、采样、训练任务、检查点、数据集、DataLoader 与输入处理器 API。 +- 所有网络失败和服务端失败均转换为携带上下文的运行时异常。 +- 通过 `TWINKLE_SERVER_URL` 与 `TWINKLE_SERVER_TOKEN` 读取默认服务地址和认证令牌。 + +## 引入依赖 + +发布到 Maven Central 后,可在项目中引入: + +```xml + + io.github.modelscope + twinkle-client-java + 1.0.0 + +``` + +在发布前,可直接将本项目导入 IntelliJ IDEA 作为 Maven 项目。 + +## 最小示例 + +```java +try (TwinkleClient client = TwinkleClient.builder() + .baseUrl(System.getenv("TWINKLE_SERVER_URL")) + .apiKey(System.getenv("TWINKLE_SERVER_TOKEN")) + .build()) { + if (!client.healthCheck()) { + throw new IllegalStateException("Twinkle 服务不可用"); + } + + var model = client.models().open("Qwen/Qwen3.6-27B"); + model.addAdapter("default", new LoraConfig(8, 16, "all-linear", 0.01, "none", null)); + model.setLoss("CrossEntropyLoss"); + model.setOptimizer("Adam", Map.of("lr", 1e-4)); +} +``` + +## 数据加载与训练 + +```java +var dataset = client.processors().dataset( + DatasetKind.DATASET, + Map.of("dataset_meta", DatasetMeta.of("ms://your-dataset"))); +dataset.setTemplate("Qwen3_5Template", Map.of("model_id", "Qwen/Qwen3.6-27B")); +dataset.encode(false, Map.of("batched", true)); + +var loader = client.processors().dataLoader(dataset.processorId(), Map.of("batch_size", 4)); +for (var batch : loader) { + model.forwardBackward(batch); + model.clipGradAndStep(1.0, 2); +} +``` + +## 配置 + +| 配置项 | 默认值 | 说明 | +| --- | --- | --- | +| `TWINKLE_SERVER_URL` | `http://127.0.0.1:8000` | 服务根地址;客户端自动补充 `/api/v1`。 | +| `TWINKLE_SERVER_TOKEN` | `EMPTY_TOKEN` | 服务端认证令牌。 | +| `routePrefix` | `/twinkle` | 会话和训练任务管理 API 的路由前缀。 | + +请勿将真实令牌、内网地址、数据集本地路径写入源码、Issue 或提交历史。 + +## 与旧版原型的迁移 + +| 原型 API | 新 API | +| --- | --- | +| `new TwinkleClient(url, token)` | `TwinkleClient.builder().baseUrl(url).apiKey(token).build()` | +| `createModel(id)` | `client.models().open(id)` | +| `createSampler(id)` | `client.samplers().open(id)` | +| `createDataset(type, args)` | `client.processors().dataset(type, args)` | +| `Map` 响应 | 稳定字段使用 record,开放字段使用 `JsonObject` / `JsonElement`。 | + +## 许可证 + +本项目采用 [Apache License 2.0](LICENSE)。 diff --git a/javaclients/SECURITY.md b/javaclients/SECURITY.md new file mode 100644 index 00000000..9f9872cf --- /dev/null +++ b/javaclients/SECURITY.md @@ -0,0 +1,5 @@ +# 安全说明 + +请不要在公开 Issue、Pull Request、日志或示例中提交 API Token、密码、内网地址和真实训练数据。 + +发现安全问题时,请使用目标 GitHub 仓库的 **Private Security Advisory** 功能进行私密报告,并说明受影响版本、复现条件和风险范围。 diff --git a/javaclients/pom.xml b/javaclients/pom.xml new file mode 100644 index 00000000..9c484c8e --- /dev/null +++ b/javaclients/pom.xml @@ -0,0 +1,65 @@ + + + 4.0.0 + io.github.modelscope + twinkle-client-java + 1.0.0-SNAPSHOT + Twinkle Java Client + Twinkle 训练服务的 Java 17 客户端 + https://github.com/modelscope/twinkle-client-java + + + Apache License 2.0 + https://www.apache.org/licenses/LICENSE-2.0.txt + + + + 17 + UTF-8 + 4.12.0 + 2.10.1 + 5.10.2 + 4.12.0 + + + + com.squareup.okhttp3 + okhttp + ${okhttp.version} + + + com.google.code.gson + gson + ${gson.version} + + + org.junit.jupiter + junit-jupiter + ${junit.version} + test + + + com.squareup.okhttp3 + mockwebserver + ${mockwebserver.version} + test + + + + + + org.apache.maven.plugins + maven-compiler-plugin + 3.13.0 + + 17 + + + + org.apache.maven.plugins + maven-surefire-plugin + 3.2.5 + + + + diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/TwinkleClient.java b/javaclients/src/main/java/io/github/modelscope/twinkle/TwinkleClient.java new file mode 100644 index 00000000..eb100268 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/TwinkleClient.java @@ -0,0 +1,151 @@ +package io.github.modelscope.twinkle; + +import io.github.modelscope.twinkle.config.ClientConfig; +import io.github.modelscope.twinkle.internal.ResponseMapper; +import io.github.modelscope.twinkle.model.ModelsClient; +import io.github.modelscope.twinkle.processor.ProcessorsClient; +import io.github.modelscope.twinkle.runs.TrainingRunsClient; +import io.github.modelscope.twinkle.sampler.SamplersClient; +import io.github.modelscope.twinkle.session.SessionManager; +import io.github.modelscope.twinkle.transport.HttpTransport; +import io.github.modelscope.twinkle.transport.OkHttpTransport; +import io.github.modelscope.twinkle.types.CapacityInfo; +import io.github.modelscope.twinkle.types.ServerCapabilities; +import java.time.Duration; +import java.util.Map; + +/** Twinkle Java SDK 的唯一顶层入口,负责配置、会话和资源客户端。 */ +public final class TwinkleClient implements AutoCloseable { + + private final HttpTransport transport; + private final String configRoutePrefix; + private final SessionManager session; + private final TrainingRunsClient trainingRuns; + private final ModelsClient models; + private final SamplersClient samplers; + private final ProcessorsClient processors; + + private TwinkleClient(Builder builder) { + ClientConfig config = builder.config.build(); + this.transport = new OkHttpTransport(config); + this.configRoutePrefix = config.routePrefix(); + this.session = new SessionManager( + transport, + config.routePrefix(), + config.heartbeatInterval(), + builder.metadata, + builder.existingSessionId + ); + this.trainingRuns = new TrainingRunsClient(transport, config.routePrefix()); + this.models = new ModelsClient(transport, config.routePrefix()); + this.samplers = new SamplersClient(transport, config.routePrefix()); + this.processors = new ProcessorsClient(transport, config.routePrefix()); + } + + public static Builder builder() { + return new Builder(); + } + + public String sessionId() { + return session.sessionId(); + } + + public TrainingRunsClient trainingRuns() { + return trainingRuns; + } + + public ModelsClient models() { + return models; + } + + public SamplersClient samplers() { + return samplers; + } + + public ProcessorsClient processors() { + return processors; + } + + public boolean healthCheck() { + try { + transport.get(configuredRoute("/healthz"), Map.of()); + return true; + } catch (RuntimeException error) { + return false; + } + } + + public ServerCapabilities serverCapabilities() { + return ResponseMapper.serverCapabilities( + transport.get(configuredRoute("/get_server_capabilities"), Map.of()).getAsJsonObject() + ); + } + + public CapacityInfo capacityInfo() { + return ResponseMapper.capacityInfo( + transport.get(configuredRoute("/capacity_info"), Map.of()).getAsJsonObject() + ); + } + + @Override + public void close() { + session.close(); + transport.close(); + } + + private String configuredRoute(String endpoint) { + return configRoutePrefix + endpoint; + } + + /** 构建客户端时可配置服务端地址、令牌和会话策略。 */ + public static final class Builder { + + private final ClientConfig.Builder config = ClientConfig.builder(); + private Map metadata = Map.of(); + private String existingSessionId; + + public Builder baseUrl(String value) { + config.baseUrl(value); + return this; + } + + public Builder apiKey(String value) { + config.apiKey(value); + return this; + } + + public Builder routePrefix(String value) { + config.routePrefix(value); + return this; + } + + public Builder connectTimeout(Duration value) { + config.connectTimeout(value); + return this; + } + + public Builder requestTimeout(Duration value) { + config.requestTimeout(value); + return this; + } + + public Builder heartbeatInterval(Duration value) { + config.heartbeatInterval(value); + return this; + } + + public Builder sessionMetadata(Map value) { + metadata = value == null ? Map.of() : Map.copyOf(value); + return this; + } + + public Builder existingSessionId(String value) { + existingSessionId = value; + return this; + } + + public TwinkleClient build() { + return new TwinkleClient(this); + } + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/config/ClientConfig.java b/javaclients/src/main/java/io/github/modelscope/twinkle/config/ClientConfig.java new file mode 100644 index 00000000..cced683b --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/config/ClientConfig.java @@ -0,0 +1,117 @@ +package io.github.modelscope.twinkle.config; + +import java.time.Duration; +import java.util.Objects; +import java.util.function.Supplier; + +/** 客户端不可变配置及其构建器。 */ +public record ClientConfig( + String apiBaseUrl, + String routePrefix, + Supplier apiKeySupplier, + Duration connectTimeout, + Duration requestTimeout, + Duration heartbeatInterval +) { + public static Builder builder() { + return new Builder(); + } + + /** 用于在构建前完成参数校验与默认值处理。 */ + public static final class Builder { + + private String baseUrl; + private String routePrefix = "/twinkle"; + private Supplier apiKeySupplier; + private Duration connectTimeout = Duration.ofSeconds(30); + private Duration requestTimeout = Duration.ofMinutes(10); + private Duration heartbeatInterval = Duration.ofSeconds(10); + + public Builder baseUrl(String value) { + this.baseUrl = value; + return this; + } + + public Builder routePrefix(String value) { + this.routePrefix = value; + return this; + } + + public Builder apiKey(String value) { + this.apiKeySupplier = () -> value; + return this; + } + + public Builder apiKeySupplier(Supplier value) { + this.apiKeySupplier = value; + return this; + } + + public Builder connectTimeout(Duration value) { + this.connectTimeout = value; + return this; + } + + public Builder requestTimeout(Duration value) { + this.requestTimeout = value; + return this; + } + + public Builder heartbeatInterval(Duration value) { + this.heartbeatInterval = value; + return this; + } + + public ClientConfig build() { + String server = normalizeBaseUrl( + baseUrl == null ? System.getenv("TWINKLE_SERVER_URL") : baseUrl + ); + String prefix = normalizePrefix(routePrefix); + Supplier supplier = apiKeySupplier == null + ? () -> System.getenv().getOrDefault("TWINKLE_SERVER_TOKEN", "EMPTY_TOKEN") + : apiKeySupplier; + validateDuration(connectTimeout, "connectTimeout"); + validateDuration(requestTimeout, "requestTimeout"); + validateDuration(heartbeatInterval, "heartbeatInterval"); + String token = Objects.requireNonNull(supplier.get(), "apiKey 不能为空").trim(); + if (token.isEmpty()) { + throw new IllegalArgumentException("apiKey 不能为空"); + } + return new ClientConfig( + server, + prefix, + supplier, + connectTimeout, + requestTimeout, + heartbeatInterval + ); + } + + private static String normalizeBaseUrl(String value) { + String result = value == null || value.isBlank() + ? "http://127.0.0.1:8000" + : value.trim(); + result = result.replaceAll("/+$", ""); + return result.endsWith("/api/v1") ? result : result + "/api/v1"; + } + + private static String normalizePrefix(String value) { + if (value == null || value.isBlank()) { + return ""; + } + String result = value.trim().replaceAll("/+$", ""); + if (!result.startsWith("/")) { + throw new IllegalArgumentException("routePrefix 必须以 / 开头"); + } + return result; + } + + private static void validateDuration(Duration value, String name) { + if ( + value == null || value.isZero() || value.isNegative() + ) { + throw new IllegalArgumentException(name + " 必须大于 0"); + } + } + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleException.java b/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleException.java new file mode 100644 index 00000000..efb24e57 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleException.java @@ -0,0 +1,13 @@ +package io.github.modelscope.twinkle.exception; + +/** Twinkle 客户端所有运行时异常的基类。 */ +public class TwinkleException extends RuntimeException { + + public TwinkleException(String message) { + super(message); + } + + public TwinkleException(String message, Throwable cause) { + super(message, cause); + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleIterationExhaustedException.java b/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleIterationExhaustedException.java new file mode 100644 index 00000000..1ba4f61f --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleIterationExhaustedException.java @@ -0,0 +1,11 @@ +package io.github.modelscope.twinkle.exception; + +import java.net.URI; + +/** HTTP 410,表示服务端远程迭代器已耗尽。 */ +public final class TwinkleIterationExhaustedException extends TwinkleServiceException { + + public TwinkleIterationExhaustedException(URI endpoint, String requestId, String detail) { + super(410, endpoint, requestId, detail); + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleProtocolException.java b/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleProtocolException.java new file mode 100644 index 00000000..f39787da --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleProtocolException.java @@ -0,0 +1,9 @@ +package io.github.modelscope.twinkle.exception; + +/** 服务端响应不是预期 JSON 协议格式。 */ +public final class TwinkleProtocolException extends TwinkleException { + + public TwinkleProtocolException(String message, Throwable cause) { + super(message, cause); + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleServiceException.java b/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleServiceException.java new file mode 100644 index 00000000..3bd46f4c --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleServiceException.java @@ -0,0 +1,36 @@ +package io.github.modelscope.twinkle.exception; + +import java.net.URI; + +/** 服务端返回非成功状态码时抛出的异常。 */ +public class TwinkleServiceException extends TwinkleException { + + private final int statusCode; + private final URI endpoint; + private final String requestId; + private final String serviceDetail; + + public TwinkleServiceException(int statusCode, URI endpoint, String requestId, String detail) { + super(detail); + this.statusCode = statusCode; + this.endpoint = endpoint; + this.requestId = requestId; + this.serviceDetail = detail; + } + + public int statusCode() { + return statusCode; + } + + public URI endpoint() { + return endpoint; + } + + public String requestId() { + return requestId; + } + + public String serviceDetail() { + return serviceDetail; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleTransportException.java b/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleTransportException.java new file mode 100644 index 00000000..7cdfcc2a --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/exception/TwinkleTransportException.java @@ -0,0 +1,9 @@ +package io.github.modelscope.twinkle.exception; + +/** 网络连接、超时或本地 HTTP 传输失败。 */ +public final class TwinkleTransportException extends TwinkleException { + + public TwinkleTransportException(String message, Throwable cause) { + super(message, cause); + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/internal/ResponseMapper.java b/javaclients/src/main/java/io/github/modelscope/twinkle/internal/ResponseMapper.java new file mode 100644 index 00000000..87e70d82 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/internal/ResponseMapper.java @@ -0,0 +1,104 @@ +package io.github.modelscope.twinkle.internal; + +import com.google.gson.JsonArray; +import com.google.gson.JsonElement; +import com.google.gson.JsonObject; +import io.github.modelscope.twinkle.exception.TwinkleProtocolException; +import io.github.modelscope.twinkle.types.CapacityInfo; +import io.github.modelscope.twinkle.types.CheckpointPath; +import io.github.modelscope.twinkle.types.DeleteCheckpointResult; +import io.github.modelscope.twinkle.types.ServerCapabilities; +import io.github.modelscope.twinkle.types.SupportedModel; +import io.github.modelscope.twinkle.types.WeightsInfo; +import java.util.ArrayList; +import java.util.List; + +/** 将服务端 JSON 响应转换为公开 record,并保留未知字段。 */ +public final class ResponseMapper { + + private ResponseMapper() {} + + public static CapacityInfo capacityInfo(JsonObject source) { + JsonObject copy = source.deepCopy(); + return new CapacityInfo( + requiredInt(copy, "max_loras"), + requiredInt(copy, "used_loras"), + requiredInt(copy, "free_loras"), + copy + ); + } + + public static ServerCapabilities serverCapabilities(JsonObject source) { + JsonObject copy = source.deepCopy(); + JsonArray values = required(copy, "supported_models").getAsJsonArray(); + List models = new ArrayList<>(); + + for (JsonElement value : values) { + JsonObject model = value.getAsJsonObject(); + String name = requiredString(model, "model_name"); + models.add(new SupportedModel(name, model)); + } + + return new ServerCapabilities(List.copyOf(models), copy); + } + + public static CheckpointPath checkpointPath( + JsonObject source, + String runId, + String checkpointId + ) { + JsonObject copy = source.deepCopy(); + String type = checkpointId.contains("/") + ? checkpointId.substring(0, checkpointId.indexOf('/')) + : ""; + + return new CheckpointPath( + requiredString(copy, "path"), + requiredString(copy, "twinkle_path"), + runId, + type, + checkpointId, + copy + ); + } + + public static DeleteCheckpointResult deleteCheckpoint(JsonObject source) { + JsonObject copy = source.deepCopy(); + boolean success = copy.has("success") && copy.remove("success").getAsBoolean(); + String message = copy.has("message") ? copy.remove("message").getAsString() : ""; + + return new DeleteCheckpointResult(success, message, copy); + } + + public static WeightsInfo weightsInfo(JsonObject source) { + JsonObject copy = source.deepCopy(); + Integer rank = copy.has("lora_rank") ? copy.remove("lora_rank").getAsInt() : null; + boolean lora = copy.has("is_lora") && copy.remove("is_lora").getAsBoolean(); + + return new WeightsInfo( + requiredString(copy, "training_run_id"), + requiredString(copy, "base_model"), + requiredString(copy, "model_owner"), + lora, + rank, + copy + ); + } + + private static int requiredInt(JsonObject value, String name) { + return required(value, name).getAsInt(); + } + + private static String requiredString(JsonObject value, String name) { + return required(value, name).getAsString(); + } + + private static JsonElement required(JsonObject value, String name) { + JsonElement result = value.remove(name); + if (result == null || result.isJsonNull()) { + throw new TwinkleProtocolException("服务端响应缺少必填字段: " + name, null); + } + + return result; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/model/ModelClient.java b/javaclients/src/main/java/io/github/modelscope/twinkle/model/ModelClient.java new file mode 100644 index 00000000..db7a59fd --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/model/ModelClient.java @@ -0,0 +1,258 @@ +package io.github.modelscope.twinkle.model; + +import com.google.gson.JsonElement; +import com.google.gson.JsonObject; +import io.github.modelscope.twinkle.transport.HttpTransport; +import io.github.modelscope.twinkle.transport.JsonValue; +import io.github.modelscope.twinkle.types.LoraConfig; +import io.github.modelscope.twinkle.types.SaveResponse; +import java.net.URLEncoder; +import java.nio.charset.StandardCharsets; +import java.util.LinkedHashMap; +import java.util.Map; + +/** 单个远程训练模型的同步操作接口。 */ +public final class ModelClient { + + private final HttpTransport transport; + private final String modelId; + private final String basePath; + private String adapterName; + + ModelClient(HttpTransport transport, String modelId) { + if (modelId == null || modelId.isBlank()) { + throw new IllegalArgumentException("modelId 不能为空"); + } + this.transport = transport; + this.modelId = stripScheme(modelId); + this.basePath = "/model/" + pathSegment(this.modelId) + "/twinkle"; + post("/create", Map.of()); + } + + public String modelId() { + return modelId; + } + + public String adapterName() { + return adapterName; + } + + public ModelClient useAdapter(String name) { + adapterName = require(name, "adapterName"); + return this; + } + + public void addAdapter(String name, LoraConfig config) { + addAdapter(name, config, Map.of()); + } + + public void addAdapter(String name, LoraConfig config, Map options) { + Map payload = new LinkedHashMap<>(); + payload.put("adapter_name", require(name, "adapterName")); + payload.put("config", config); + payload.putAll(options == null ? Map.of() : options); + post("/add_adapter_to_model", payload); + adapterName = name; + } + + public JsonElement forward(JsonElement inputs) { + return result(post("/forward", withAdapter(Map.of("inputs", new JsonValue(inputs))))); + } + + public JsonElement forwardOnly(JsonElement inputs) { + return result(post("/forward_only", withAdapter(Map.of("inputs", new JsonValue(inputs))))); + } + + public JsonElement forwardBackward(JsonElement inputs) { + return result( + post("/forward_backward", withAdapter(Map.of("inputs", new JsonValue(inputs)))) + ); + } + + public void backward() { + post("/backward", withAdapter(Map.of())); + } + + public double calculateLoss() { + return result(post("/calculate_loss", withAdapter(Map.of()))).getAsDouble(); + } + + public JsonObject calculateMetric(boolean training) { + return result( + post("/calculate_metric", withAdapter(Map.of("is_training", training))) + ).getAsJsonObject(); + } + + public void setLoss(String lossClass) { + post("/set_loss", withAdapter(Map.of("loss_cls", require(lossClass, "lossClass")))); + } + + public void setOptimizer(String optimizerClass, Map options) { + post( + "/set_optimizer", + withAdapter( + merge(Map.of("optimizer_cls", require(optimizerClass, "optimizerClass")), options) + ) + ); + } + + public void setLrScheduler(String schedulerClass, Map options) { + post( + "/set_lr_scheduler", + withAdapter( + merge(Map.of("scheduler_cls", require(schedulerClass, "schedulerClass")), options) + ) + ); + } + + public void step() { + post("/step", withAdapter(Map.of())); + } + + public void zeroGrad() { + post("/zero_grad", withAdapter(Map.of())); + } + + public void lrStep() { + post("/lr_step", withAdapter(Map.of())); + } + + public String clipGradNorm(double maxGradNorm, int normType) { + return result( + post( + "/clip_grad_norm", + withAdapter(Map.of("max_grad_norm", maxGradNorm, "norm_type", normType)) + ) + ).getAsString(); + } + + public void clipGradAndStep(double maxGradNorm, int normType) { + post( + "/clip_grad_and_step", + withAdapter(Map.of("max_grad_norm", maxGradNorm, "norm_type", normType)) + ); + } + + public void setTemplate(String templateClass, Map options) { + post( + "/set_template", + withAdapter( + merge( + Map.of( + "template_cls", + require(templateClass, "templateClass"), + "model_id", + modelId + ), + options + ) + ) + ); + } + + public void setProcessor(String processorClass, Map options) { + post( + "/set_processor", + withAdapter( + merge(Map.of("processor_cls", require(processorClass, "processorClass")), options) + ) + ); + } + + public void addMetric(String metricClass, Boolean training) { + Map data = new LinkedHashMap<>(); + data.put("metric_cls", require(metricClass, "metricClass")); + if (training != null) { + data.put("is_training", training); + } + post("/add_metric", withAdapter(data)); + } + + public void applyPatch(String patchClass) { + post("/apply_patch", withAdapter(Map.of("patch_cls", require(patchClass, "patchClass")))); + } + + public JsonObject stateDict() { + return result(post("/get_state_dict", withAdapter(Map.of()))).getAsJsonObject(); + } + + public String trainConfigs() { + return result(post("/get_train_configs", withAdapter(Map.of()))).getAsString(); + } + + public SaveResponse save(String name, boolean saveOptimizer) { + JsonObject value = post( + "/save", + withAdapter(Map.of("name", require(name, "name"), "save_optimizer", saveOptimizer)) + ).getAsJsonObject(); + return new SaveResponse(string(value, "twinkle_path"), string(value, "checkpoint_dir")); + } + + public void load(String name, boolean loadOptimizer) { + post( + "/load", + withAdapter(Map.of("name", require(name, "name"), "load_optimizer", loadOptimizer)) + ); + } + + public JsonObject resumeFromCheckpoint(String name, boolean resumeOnlyModel) { + return result( + post( + "/resume_from_checkpoint", + withAdapter( + Map.of("name", require(name, "name"), "resume_only_model", resumeOnlyModel) + ) + ) + ).getAsJsonObject(); + } + + private JsonElement post(String endpoint, Map payload) { + return transport.post(basePath + endpoint, payload); + } + + private Map withAdapter(Map values) { + Map result = new LinkedHashMap<>(); + result.putAll(values); + if (adapterName != null) { + result.put("adapter_name", adapterName); + } + return result; + } + + private static Map merge(Map first, Map second) { + Map result = new LinkedHashMap<>(); + result.putAll(first); + if (second != null) { + result.putAll(second); + } + return result; + } + + private static JsonElement result(JsonElement value) { + return value.isJsonObject() && value.getAsJsonObject().has("result") + ? value.getAsJsonObject().get("result") + : value; + } + + private static String stripScheme(String value) { + int index = value.indexOf("://"); + return index >= 0 ? value.substring(index + 3) : value; + } + + private static String pathSegment(String value) { + return URLEncoder.encode(value, StandardCharsets.UTF_8).replace("+", "%20"); + } + + private static String require(String value, String name) { + if (value == null || value.isBlank()) { + throw new IllegalArgumentException(name + " 不能为空"); + } + return value; + } + + private static String string(JsonObject value, String name) { + return value.has(name) && !value.get(name).isJsonNull() + ? value.get(name).getAsString() + : null; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/model/ModelsClient.java b/javaclients/src/main/java/io/github/modelscope/twinkle/model/ModelsClient.java new file mode 100644 index 00000000..b46273d8 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/model/ModelsClient.java @@ -0,0 +1,17 @@ +package io.github.modelscope.twinkle.model; + +import io.github.modelscope.twinkle.transport.HttpTransport; + +/** 用于打开服务端训练模型的工厂。 */ +public final class ModelsClient { + + private final HttpTransport transport; + + public ModelsClient(HttpTransport transport, String ignoredRoutePrefix) { + this.transport = transport; + } + + public ModelClient open(String modelId) { + return new ModelClient(transport, modelId); + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/processor/DataLoaderClient.java b/javaclients/src/main/java/io/github/modelscope/twinkle/processor/DataLoaderClient.java new file mode 100644 index 00000000..3bc8c98a --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/processor/DataLoaderClient.java @@ -0,0 +1,67 @@ +package io.github.modelscope.twinkle.processor; + +import com.google.gson.JsonElement; +import io.github.modelscope.twinkle.exception.TwinkleIterationExhaustedException; +import io.github.modelscope.twinkle.transport.HttpTransport; +import java.util.Iterator; +import java.util.Map; +import java.util.NoSuchElementException; + +/** 支持 Java for-each 的远程数据加载器。 */ +public final class DataLoaderClient extends RemoteProcessor implements Iterable { + + DataLoaderClient(HttpTransport transport, String id) { + super(transport, id); + } + + public int length() { + return call("__len__", Map.of()).getAsInt(); + } + + public JsonElement setProcessor(String processorClass, Map options) { + return call( + "set_processor", + options == null + ? Map.of("processor_cls", processorClass) + : merge(processorClass, options) + ); + } + + public JsonElement skipConsumedSamples(int count) { + return call("skip_consumed_samples", Map.of("consumed_train_samples", count)); + } + + public JsonElement state() { + return call("get_state", Map.of()); + } + + @Override + public Iterator iterator() { + call("__iter__", Map.of()); + return new Iterator<>() { + private boolean exhausted; + + @Override + public boolean hasNext() { + return !exhausted; + } + + @Override + public JsonElement next() { + try { + return call("__next__", Map.of()); + } catch (TwinkleIterationExhaustedException error) { + exhausted = true; + throw new NoSuchElementException(error.getMessage()); + } + } + }; + } + + private static Map merge(String processorClass, Map options) { + java.util.LinkedHashMap result = new java.util.LinkedHashMap<>(); + result.put("processor_cls", processorClass); + result.putAll(options); + return result; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/processor/DatasetClient.java b/javaclients/src/main/java/io/github/modelscope/twinkle/processor/DatasetClient.java new file mode 100644 index 00000000..6d1e6542 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/processor/DatasetClient.java @@ -0,0 +1,138 @@ +package io.github.modelscope.twinkle.processor; + +import com.google.gson.JsonElement; +import io.github.modelscope.twinkle.transport.HttpTransport; +import io.github.modelscope.twinkle.types.DatasetKind; +import io.github.modelscope.twinkle.types.DatasetMeta; +import java.util.LinkedHashMap; +import java.util.Map; + +/** 远程数据集操作接口。 */ +public final class DatasetClient extends RemoteProcessor { + + private final DatasetKind kind; + + DatasetClient(HttpTransport transport, String id, DatasetKind kind) { + super(transport, id); + this.kind = kind; + } + + public DatasetKind kind() { + return kind; + } + + public JsonElement setTemplate(String templateFunction, Map options) { + return call("set_template", merge(values("template_func", templateFunction), options)); + } + + public JsonElement encode(boolean addGenerationPrompt, Map options) { + return call("encode", merge(Map.of("add_generation_prompt", addGenerationPrompt), options)); + } + + public JsonElement check(Map options) { + return call("check", options); + } + + public JsonElement castColumn(String column, boolean decode) { + return call("cast_column", Map.of("column", column, "decode", decode)); + } + + public JsonElement map( + String function, + DatasetMeta meta, + Map initArgs, + Map options + ) { + return call( + "map", + merge( + values("preprocess_func", function, "dataset_meta", meta, "init_args", initArgs), + options + ) + ); + } + + public JsonElement filter( + String function, + DatasetMeta meta, + Map initArgs, + Map options + ) { + return call( + "filter", + merge( + values("filter_func", function, "dataset_meta", meta, "init_args", initArgs), + options + ) + ); + } + + public JsonElement addDataset(DatasetMeta meta, Map options) { + return call("add_dataset", merge(Map.of("dataset_meta", meta), options)); + } + + public JsonElement mixDataset(boolean interleave) { + return call("mix_dataset", Map.of("interleave", interleave)); + } + + public JsonElement saveAs( + String outputPath, + String format, + int batchSize, + String mode, + Map options + ) { + return call( + "save_as", + merge( + values( + "output_path", + outputPath, + "format", + format, + "batch_size", + batchSize, + "mode", + mode + ), + options + ) + ); + } + + public JsonElement flushSave() { + return call("flush_save", Map.of()); + } + + public JsonElement getItem(int index) { + return call("__getitem__", Map.of("idx", index)); + } + + public int length() { + return call("__len__", Map.of()).getAsInt(); + } + + public JsonElement packDataset() { + return call("pack_dataset", Map.of()); + } + + private static Map merge(Map first, Map second) { + Map result = new LinkedHashMap<>(); + result.putAll(first); + if (second != null) { + result.putAll(second); + } + return result; + } + + private static Map values(Object... pairs) { + Map result = new LinkedHashMap<>(); + for (int index = 0; index < pairs.length; index += 2) { + result.put( + (String) pairs[index], + pairs[index + 1] + ); + } + return result; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/processor/InputProcessorClient.java b/javaclients/src/main/java/io/github/modelscope/twinkle/processor/InputProcessorClient.java new file mode 100644 index 00000000..97266242 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/processor/InputProcessorClient.java @@ -0,0 +1,24 @@ +package io.github.modelscope.twinkle.processor; + +import com.google.gson.JsonElement; +import io.github.modelscope.twinkle.transport.HttpTransport; +import io.github.modelscope.twinkle.transport.JsonValue; +import java.util.List; +import java.util.Map; + +/** 远程输入处理器接口。 */ +public final class InputProcessorClient extends RemoteProcessor { + + InputProcessorClient(HttpTransport transport, String id) { + super(transport, id); + } + + public JsonElement process(List inputs, Map options) { + java.util.LinkedHashMap data = new java.util.LinkedHashMap<>(); + data.put("inputs", inputs.stream().map(JsonValue::new).toList()); + if (options != null) { + data.putAll(options); + } + return call("__call__", data); + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/processor/ProcessorsClient.java b/javaclients/src/main/java/io/github/modelscope/twinkle/processor/ProcessorsClient.java new file mode 100644 index 00000000..0168fab0 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/processor/ProcessorsClient.java @@ -0,0 +1,53 @@ +package io.github.modelscope.twinkle.processor; + +import com.google.gson.JsonObject; +import io.github.modelscope.twinkle.transport.HttpTransport; +import io.github.modelscope.twinkle.types.DatasetKind; +import java.util.LinkedHashMap; +import java.util.Map; + +/** 创建远程数据集、加载器和输入处理器。 */ +public final class ProcessorsClient { + + private final HttpTransport transport; + + public ProcessorsClient(HttpTransport transport, String ignoredRoutePrefix) { + this.transport = transport; + } + + public DatasetClient dataset(DatasetKind kind, Map options) { + Map data = new LinkedHashMap<>(); + data.put("processor_type", "dataset"); + data.put("class_type", kind.serverClassName()); + if (options != null) { + data.putAll(options); + } + return new DatasetClient(transport, create(data), kind); + } + + public DataLoaderClient dataLoader(String datasetProcessorId, Map options) { + Map data = new LinkedHashMap<>(); + data.put("processor_type", "dataloader"); + data.put("class_type", "DataLoader"); + data.put("dataset", datasetProcessorId); + if (options != null) { + data.putAll(options); + } + return new DataLoaderClient(transport, create(data)); + } + + public InputProcessorClient inputProcessor(Map options) { + Map data = new LinkedHashMap<>(); + data.put("processor_type", "processor"); + data.put("class_type", "InputProcessor"); + if (options != null) { + data.putAll(options); + } + return new InputProcessorClient(transport, create(data)); + } + + private String create(Map data) { + JsonObject response = transport.post("/processor/twinkle/create", data).getAsJsonObject(); + return response.get("processor_id").getAsString(); + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/processor/RemoteProcessor.java b/javaclients/src/main/java/io/github/modelscope/twinkle/processor/RemoteProcessor.java new file mode 100644 index 00000000..a9520360 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/processor/RemoteProcessor.java @@ -0,0 +1,34 @@ +package io.github.modelscope.twinkle.processor; + +import com.google.gson.JsonElement; +import com.google.gson.JsonObject; +import io.github.modelscope.twinkle.transport.HttpTransport; +import java.util.LinkedHashMap; +import java.util.Map; + +/** 所有远程处理器共享的调用封装。 */ +abstract class RemoteProcessor { + + protected final HttpTransport transport; + protected final String processorId; + + RemoteProcessor(HttpTransport transport, String processorId) { + this.transport = transport; + this.processorId = processorId; + } + + public String processorId() { + return processorId; + } + + protected JsonElement call(String function, Map arguments) { + Map payload = new LinkedHashMap<>(); + payload.put("processor_id", processorId); + payload.put("function", function); + if (arguments != null) { + payload.putAll(arguments); + } + JsonObject response = transport.post("/processor/twinkle/call", payload).getAsJsonObject(); + return response.has("result") ? response.get("result") : response; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/runs/TrainingRunsClient.java b/javaclients/src/main/java/io/github/modelscope/twinkle/runs/TrainingRunsClient.java new file mode 100644 index 00000000..2a9b8908 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/runs/TrainingRunsClient.java @@ -0,0 +1,160 @@ +package io.github.modelscope.twinkle.runs; + +import com.google.gson.JsonArray; +import com.google.gson.JsonElement; +import com.google.gson.JsonObject; +import io.github.modelscope.twinkle.internal.ResponseMapper; +import io.github.modelscope.twinkle.transport.HttpTransport; +import io.github.modelscope.twinkle.types.Checkpoint; +import io.github.modelscope.twinkle.types.CheckpointPath; +import io.github.modelscope.twinkle.types.Cursor; +import io.github.modelscope.twinkle.types.DeleteCheckpointResult; +import io.github.modelscope.twinkle.types.TrainingRun; +import io.github.modelscope.twinkle.types.WeightsInfo; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; + +/** 训练任务和检查点查询接口。 */ +public final class TrainingRunsClient { + + private final HttpTransport transport; + private final String prefix; + + public TrainingRunsClient(HttpTransport transport, String prefix) { + this.transport = transport; + this.prefix = prefix; + } + + public Page list(int limit, int offset, boolean allUsers) { + if (limit <= 0 || offset < 0) { + throw new IllegalArgumentException("limit 必须大于 0,offset 不能小于 0"); + } + JsonObject body = transport + .get( + prefix + "/training_runs", + Map.of("limit", limit, "offset", offset, "all_users", allUsers) + ) + .getAsJsonObject(); + List runs = new ArrayList<>(); + for (JsonElement item : body.getAsJsonArray("training_runs")) { + runs.add(toRun(item.getAsJsonObject())); + } + + JsonObject cursor = body.has("cursor") + ? body.getAsJsonObject("cursor") + : new JsonObject(); + return new Page( + List.copyOf(runs), + new Cursor( + intValue(cursor, "limit"), + intValue(cursor, "offset"), + intValue(cursor, "total_count") + ) + ); + } + + public TrainingRun get(String runId) { + return toRun( + transport.get(prefix + "/training_runs/" + requireId(runId), Map.of()).getAsJsonObject() + ); + } + + public List listCheckpoints(String runId) { + JsonArray values = transport + .get(prefix + "/training_runs/" + requireId(runId) + "/checkpoints", Map.of()) + .getAsJsonObject() + .getAsJsonArray("checkpoints"); + List result = new ArrayList<>(); + for (JsonElement item : values) { + result.add(toCheckpoint(item.getAsJsonObject())); + } + return List.copyOf(result); + } + + public CheckpointPath checkpointPath(String runId, String checkpointId) { + return ResponseMapper.checkpointPath( + transport + .get( + prefix + "/checkpoint_path/" + requireId(runId) + "/" + requireId(checkpointId), + Map.of() + ) + .getAsJsonObject(), + runId, + checkpointId + ); + } + + public DeleteCheckpointResult deleteCheckpoint(String runId, String checkpointId) { + return ResponseMapper.deleteCheckpoint( + transport + .delete( + prefix + + "/training_runs/" + + requireId(runId) + + "/checkpoints/" + + requireId(checkpointId) + ) + .getAsJsonObject() + ); + } + + public WeightsInfo weightsInfo(String twinklePath) { + return ResponseMapper.weightsInfo( + transport + .post(prefix + "/weights_info", Map.of("twinkle_path", requireId(twinklePath))) + .getAsJsonObject() + ); + } + + public String latestCheckpointPath(String runId) { + List checkpoints = listCheckpoints(runId); + return checkpoints.isEmpty() + ? null + : checkpointPath(runId, checkpoints.get(checkpoints.size() - 1).checkpointId()).path(); + } + + private static TrainingRun toRun(JsonObject value) { + return new TrainingRun( + string(value, "training_run_id"), + string(value, "base_model"), + string(value, "model_owner"), + value + ); + } + + private static Checkpoint toCheckpoint(JsonObject value) { + return new Checkpoint( + string(value, "checkpoint_id"), + string(value, "checkpoint_type"), + string(value, "twinkle_path"), + value + ); + } + + private static String string(JsonObject value, String key) { + return value.has(key) && !value.get(key).isJsonNull() + ? value.get(key).getAsString() + : null; + } + + private static int intValue(JsonObject value, String key) { + return value.has(key) + ? value.get(key).getAsInt() + : 0; + } + + private static String requireId(String value) { + if (value == null || value.isBlank()) { + throw new IllegalArgumentException("标识符不能为空"); + } + return value.replace("/", "%2F"); + } + + /** 训练任务分页结果。 */ + public record Page( + List runs, + Cursor cursor + ) { + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/sampler/SamplerClient.java b/javaclients/src/main/java/io/github/modelscope/twinkle/sampler/SamplerClient.java new file mode 100644 index 00000000..c25563c4 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/sampler/SamplerClient.java @@ -0,0 +1,124 @@ +package io.github.modelscope.twinkle.sampler; + +import com.google.gson.JsonArray; +import com.google.gson.JsonElement; +import com.google.gson.JsonObject; +import io.github.modelscope.twinkle.transport.HttpTransport; +import io.github.modelscope.twinkle.transport.JsonValue; +import io.github.modelscope.twinkle.types.LoraConfig; +import io.github.modelscope.twinkle.types.SampleRequest; +import io.github.modelscope.twinkle.types.SampleResult; +import io.github.modelscope.twinkle.types.SampledSequence; +import java.net.URLEncoder; +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +/** 单个远程采样器的同步操作接口。 */ +public final class SamplerClient { + + private final HttpTransport transport; + private final String basePath; + private String adapterName; + + SamplerClient(HttpTransport transport, String modelId) { + if (modelId == null || modelId.isBlank()) { + throw new IllegalArgumentException("modelId 不能为空"); + } + this.transport = transport; + String normalized = stripScheme(modelId); + this.basePath = + "/sampler/" + + URLEncoder.encode(normalized, StandardCharsets.UTF_8).replace("+", "%20") + + "/twinkle"; + post("/create", Map.of()); + } + + public SamplerClient useAdapter(String value) { + if (value == null || value.isBlank()) { + throw new IllegalArgumentException("adapterName 不能为空"); + } + adapterName = value; + return this; + } + + public JsonObject addAdapter(String name, LoraConfig config) { + JsonObject result = post( + "/add_adapter_to_sampler", + Map.of("adapter_name", name, "config", config) + ).getAsJsonObject(); + adapterName = name; + return result; + } + + public List sample(SampleRequest request) { + Map data = new LinkedHashMap<>(); + data.put("inputs", request.inputs().stream().map(JsonValue::new).toList()); + data.put("sampling_params", request.samplingParams()); + data.put("adapter_name", request.adapterName()); + data.put("num_samples", request.numSamples()); + if (request.adapterUri() != null) { + data.put("adapter_uri", request.adapterUri()); + } + JsonArray samples = post("/sample", data).getAsJsonObject().getAsJsonArray("samples"); + List results = new ArrayList<>(); + for (JsonElement sample : samples) { + results.add(parseResult(sample.getAsJsonObject())); + } + return List.copyOf(results); + } + + public void setTemplate(String templateClass, String adapter, Map options) { + Map data = new LinkedHashMap<>(); + data.put("template_cls", templateClass); + data.put("adapter_name", adapter == null ? "" : adapter); + if (options != null) { + data.putAll(options); + } + post("/set_template", data); + } + + public void applyPatch(String patchClass) { + post( + "/apply_patch", + Map.of("patch_cls", patchClass, "adapter_name", adapterName == null ? "" : adapterName) + ); + } + + private JsonElement post(String endpoint, Map data) { + return transport.post(basePath + endpoint, data); + } + + private static SampleResult parseResult(JsonObject value) { + List sequences = new ArrayList<>(); + for (JsonElement item : value.getAsJsonArray("sequences")) { + JsonObject sequence = item.getAsJsonObject(); + List tokens = new ArrayList<>(); + for (JsonElement token : sequence.getAsJsonArray("tokens")) { + tokens.add(token.getAsInt()); + } + sequences.add( + new SampledSequence( + string(sequence, "stop_reason"), + List.copyOf(tokens), + string(sequence, "decoded"), + sequence + ) + ); + } + return new SampleResult(List.copyOf(sequences)); + } + + private static String string(JsonObject value, String key) { + return value.has(key) && !value.get(key).isJsonNull() + ? value.get(key).getAsString() + : null; + } + + private static String stripScheme(String value) { + int index = value.indexOf("://"); + return index >= 0 ? value.substring(index + 3) : value; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/sampler/SamplersClient.java b/javaclients/src/main/java/io/github/modelscope/twinkle/sampler/SamplersClient.java new file mode 100644 index 00000000..0896c108 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/sampler/SamplersClient.java @@ -0,0 +1,17 @@ +package io.github.modelscope.twinkle.sampler; + +import io.github.modelscope.twinkle.transport.HttpTransport; + +/** 用于打开服务端采样器的工厂。 */ +public final class SamplersClient { + + private final HttpTransport transport; + + public SamplersClient(HttpTransport transport, String ignoredRoutePrefix) { + this.transport = transport; + } + + public SamplerClient open(String modelId) { + return new SamplerClient(transport, modelId); + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/session/SessionManager.java b/javaclients/src/main/java/io/github/modelscope/twinkle/session/SessionManager.java new file mode 100644 index 00000000..0f527f7f --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/session/SessionManager.java @@ -0,0 +1,68 @@ +package io.github.modelscope.twinkle.session; + +import com.google.gson.JsonElement; +import io.github.modelscope.twinkle.transport.HttpTransport; +import java.time.Duration; +import java.util.Map; +import java.util.concurrent.Executors; +import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.TimeUnit; +import java.util.logging.Level; +import java.util.logging.Logger; + +/** 创建会话并管理后台心跳,关闭时不会再发起心跳。 */ +public final class SessionManager implements AutoCloseable { + + private static final Logger LOG = Logger.getLogger(SessionManager.class.getName()); + private final HttpTransport transport; + private final String routePrefix; + private final ScheduledExecutorService executor; + private final String sessionId; + + public SessionManager( + HttpTransport transport, + String routePrefix, + Duration interval, + Map metadata, + String existingSessionId + ) { + this.transport = transport; + this.routePrefix = routePrefix; + this.sessionId = existingSessionId == null || existingSessionId.isBlank() + ? create(metadata) + : existingSessionId; + transport.setSessionId(sessionId); + this.executor = Executors.newSingleThreadScheduledExecutor(runnable -> { + Thread thread = new Thread(runnable, "TwinkleSessionHeartbeat"); + thread.setDaemon(true); + return thread; + }); + long delay = interval.toMillis(); + executor.scheduleWithFixedDelay(this::heartbeat, delay, delay, TimeUnit.MILLISECONDS); + } + + public String sessionId() { + return sessionId; + } + + private String create(Map metadata) { + JsonElement response = transport.post( + routePrefix + "/create_session", + Map.of("metadata", metadata) + ); + return response.getAsJsonObject().get("session_id").getAsString(); + } + + private void heartbeat() { + try { + transport.post(routePrefix + "/session_heartbeat", Map.of("session_id", sessionId)); + } catch (RuntimeException error) { + LOG.log(Level.WARNING, "Twinkle 会话心跳失败", error); + } + } + + @Override + public void close() { + executor.shutdownNow(); + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/transport/HttpTransport.java b/javaclients/src/main/java/io/github/modelscope/twinkle/transport/HttpTransport.java new file mode 100644 index 00000000..503f57f0 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/transport/HttpTransport.java @@ -0,0 +1,16 @@ +package io.github.modelscope.twinkle.transport; + +import com.google.gson.JsonElement; +import java.util.Map; + +/** 资源客户端使用的最小 HTTP 抽象。 */ +public interface HttpTransport extends AutoCloseable { + JsonElement get(String path, Map query); + JsonElement post(String path, Map payload); + JsonElement delete(String path); + void setSessionId(String sessionId); + String sessionId(); + + @Override + void close(); +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/transport/JsonValue.java b/javaclients/src/main/java/io/github/modelscope/twinkle/transport/JsonValue.java new file mode 100644 index 00000000..344f533f --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/transport/JsonValue.java @@ -0,0 +1,11 @@ +package io.github.modelscope.twinkle.transport; + +import com.google.gson.JsonElement; +import java.util.Objects; + +/** 显式传递原始 JSON 的包装类型,避免被再次转换为字符串。 */ +public record JsonValue(JsonElement value) { + public JsonValue { + Objects.requireNonNull(value, "value 不能为空"); + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/transport/OkHttpTransport.java b/javaclients/src/main/java/io/github/modelscope/twinkle/transport/OkHttpTransport.java new file mode 100644 index 00000000..716c07c3 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/transport/OkHttpTransport.java @@ -0,0 +1,147 @@ +package io.github.modelscope.twinkle.transport; + +import com.google.gson.JsonElement; +import com.google.gson.JsonObject; +import com.google.gson.JsonParser; +import io.github.modelscope.twinkle.config.ClientConfig; +import io.github.modelscope.twinkle.exception.TwinkleIterationExhaustedException; +import io.github.modelscope.twinkle.exception.TwinkleProtocolException; +import io.github.modelscope.twinkle.exception.TwinkleServiceException; +import io.github.modelscope.twinkle.exception.TwinkleTransportException; +import java.io.IOException; +import java.net.URI; +import java.time.Duration; +import java.util.Map; +import java.util.UUID; +import okhttp3.HttpUrl; +import okhttp3.MediaType; +import okhttp3.OkHttpClient; +import okhttp3.Request; +import okhttp3.RequestBody; +import okhttp3.Response; + +/** 基于 OkHttp 的同步传输实现,集中处理请求头和异常。 */ +public final class OkHttpTransport implements HttpTransport { + + private static final MediaType JSON = MediaType.get("application/json; charset=utf-8"); + private final ClientConfig config; + private final OkHttpClient client; + private final TwinkleJsonCodec codec = new TwinkleJsonCodec(); + private final String requestId = UUID.randomUUID().toString(); + private volatile String sessionId; + + public OkHttpTransport(ClientConfig config) { + this.config = config; + this.client = new OkHttpClient.Builder() + .connectTimeout(config.connectTimeout()) + .readTimeout(config.requestTimeout()) + .writeTimeout(config.requestTimeout()) + .callTimeout(config.requestTimeout()) + .build(); + } + + @Override + public JsonElement get(String path, Map query) { + HttpUrl.Builder url = url(path).newBuilder(); + if (query != null) { + query.forEach((key, value) -> { + if (value != null) { + url.addQueryParameter(key, String.valueOf(value)); + } + }); + } + return execute(new Request.Builder().url(url.build()).get()); + } + + @Override + public JsonElement post(String path, Map payload) { + String body = codec.gson().toJson(codec.encode(payload)); + return execute(new Request.Builder().url(url(path)).post(RequestBody.create(body, JSON))); + } + + @Override + public JsonElement delete(String path) { + return execute(new Request.Builder().url(url(path)).delete()); + } + + @Override + public void setSessionId(String sessionId) { + this.sessionId = sessionId; + } + + @Override + public String sessionId() { + return sessionId; + } + + @Override + public void close() { + client.dispatcher().executorService().shutdown(); + client.connectionPool().evictAll(); + } + + private HttpUrl url(String path) { + String normalized = path.startsWith("/") ? path : "/" + path; + HttpUrl parsed = HttpUrl.parse(config.apiBaseUrl() + normalized); + if (parsed == null) { + throw new IllegalArgumentException("无效的请求地址: " + path); + } + return parsed; + } + + private JsonElement execute(Request.Builder request) { + String token = config.apiKeySupplier().get(); + if (token == null || token.isBlank()) { + throw new IllegalStateException("apiKey 不能为空"); + } + String authorization = "Bearer " + token; + request + .header("Authorization", authorization) + .header("Twinkle-Authorization", authorization) + .header("x-request-id", requestId) + .header("X-Ray-Serve-Request-Id", requestId) + .header("serve_multiplexed_model_id", requestId) + .header("Serve-Multiplexed-Model-Id", requestId); + if (sessionId != null && !sessionId.isBlank()) { + request.header("X-Twinkle-Session-Id", sessionId); + } + try (Response response = client.newCall(request.build()).execute()) { + String text = response.body() == null ? "" : response.body().string(); + URI endpoint = response.request().url().uri(); + if (!response.isSuccessful()) { + String detail = detail(text); + if (response.code() == 410) { + throw new TwinkleIterationExhaustedException(endpoint, requestId, detail); + } + throw new TwinkleServiceException(response.code(), endpoint, requestId, detail); + } + try { + return text.isBlank() ? new JsonObject() : JsonParser.parseString(text); + } catch (RuntimeException error) { + throw new TwinkleProtocolException( + "服务端响应不是合法 JSON: " + endpoint, + error + ); + } + } catch (TwinkleServiceException | TwinkleProtocolException error) { + throw error; + } catch (IOException error) { + throw new TwinkleTransportException( + "HTTP 请求失败: " + request.build().url(), + error + ); + } + } + + private String detail(String text) { + try { + JsonElement json = JsonParser.parseString(text); + if (json.isJsonObject() && json.getAsJsonObject().has("detail")) { + return json.getAsJsonObject().get("detail").getAsString(); + } + } catch (RuntimeException ignored) { + /* 响应正文不是 JSON 时直接使用原文。 */ + } + return text.isBlank() ? "服务端未返回错误详情" : text; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/transport/TwinkleJsonCodec.java b/javaclients/src/main/java/io/github/modelscope/twinkle/transport/TwinkleJsonCodec.java new file mode 100644 index 00000000..23c8f265 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/transport/TwinkleJsonCodec.java @@ -0,0 +1,50 @@ +package io.github.modelscope.twinkle.transport; + +import com.google.gson.Gson; +import com.google.gson.GsonBuilder; +import com.google.gson.JsonArray; +import com.google.gson.JsonElement; +import com.google.gson.JsonNull; +import com.google.gson.JsonObject; +import com.google.gson.JsonPrimitive; +import java.util.Collection; +import java.util.Map; + +/** 负责递归编码 Twinkle 请求负载。 */ +public final class TwinkleJsonCodec { + + private final Gson gson = new GsonBuilder().serializeNulls().create(); + + public Gson gson() { + return gson; + } + + public JsonElement encode(Object value) { + if (value == null) { + return JsonNull.INSTANCE; + } + if (value instanceof JsonValue raw) { + return raw.value(); + } + if (value instanceof JsonElement json) { + return json; + } + if (value instanceof TwinkleSerializable serializable) { + return new JsonPrimitive(gson.toJson(serializable.toTwinkleJson())); + } + if (value instanceof Map map) { + JsonObject object = new JsonObject(); + map.forEach((key, item) -> object.add(String.valueOf(key), encode(item))); + return object; + } + if (value instanceof Collection collection) { + JsonArray array = new JsonArray(); + collection.forEach(item -> array.add(encode(item))); + return array; + } + if (value.getClass().isArray()) { + return gson.toJsonTree(value); + } + return gson.toJsonTree(value); + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/transport/TwinkleSerializable.java b/javaclients/src/main/java/io/github/modelscope/twinkle/transport/TwinkleSerializable.java new file mode 100644 index 00000000..e098140e --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/transport/TwinkleSerializable.java @@ -0,0 +1,8 @@ +package io.github.modelscope.twinkle.transport; + +import com.google.gson.JsonObject; + +/** 需要按 Twinkle 特殊 JSON 字符串协议传输的值对象。 */ +public interface TwinkleSerializable { + JsonObject toTwinkleJson(); +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/CapacityInfo.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/CapacityInfo.java new file mode 100644 index 00000000..4c40dd87 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/CapacityInfo.java @@ -0,0 +1,12 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonObject; + +/** 服务端 LoRA 容量信息。 */ +public record CapacityInfo( + int maxLoras, + int usedLoras, + int freeLoras, + JsonObject extensions +) { +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/Checkpoint.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/Checkpoint.java new file mode 100644 index 00000000..3e18e35d --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/Checkpoint.java @@ -0,0 +1,11 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonObject; + +/** 检查点摘要;服务端新增字段保留在 raw 中。 */ +public record Checkpoint( + String checkpointId, + String checkpointType, + String twinklePath, + JsonObject raw +) {} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/CheckpointPage.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/CheckpointPage.java new file mode 100644 index 00000000..e35779fb --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/CheckpointPage.java @@ -0,0 +1,12 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonObject; +import java.util.List; + +/** 检查点分页响应。 */ +public record CheckpointPage( + List checkpoints, + Cursor cursor, + JsonObject extensions +) { +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/CheckpointPath.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/CheckpointPath.java new file mode 100644 index 00000000..5ffe2cb1 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/CheckpointPath.java @@ -0,0 +1,13 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonObject; + +/** 检查点标识符对应的本地路径和 Twinkle 路径。 */ +public record CheckpointPath( + String path, + String twinklePath, + String trainingRunId, + String checkpointType, + String checkpointId, + JsonObject extensions +) {} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/Cursor.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/Cursor.java new file mode 100644 index 00000000..97ce74d2 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/Cursor.java @@ -0,0 +1,9 @@ +package io.github.modelscope.twinkle.types; + +/** 列表接口返回的分页游标。 */ +public record Cursor( + int limit, + int offset, + int totalCount +) { +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/DatasetKind.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/DatasetKind.java new file mode 100644 index 00000000..653d7be6 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/DatasetKind.java @@ -0,0 +1,20 @@ +package io.github.modelscope.twinkle.types; + +/** 与服务端类名对应的数据集类型。 */ +public enum DatasetKind { + DATASET("Dataset"), + LAZY_DATASET("LazyDataset"), + ITERABLE_DATASET("IterableDataset"), + PACKING_DATASET("PackingDataset"), + ITERABLE_PACKING_DATASET("IterablePackingDataset"); + + private final String serverClassName; + + DatasetKind(String serverClassName) { + this.serverClassName = serverClassName; + } + + public String serverClassName() { + return serverClassName; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/DatasetMeta.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/DatasetMeta.java new file mode 100644 index 00000000..cff47790 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/DatasetMeta.java @@ -0,0 +1,74 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.Gson; +import com.google.gson.JsonArray; +import com.google.gson.JsonObject; +import io.github.modelscope.twinkle.transport.TwinkleSerializable; +import java.util.List; + +/** 远程数据集的定位信息。 */ +public record DatasetMeta( + String datasetId, + String subsetName, + String split, + Object dataSlice, + Object data +) + implements TwinkleSerializable { + public DatasetMeta { + if ( + (datasetId == null || datasetId.isBlank()) && data == null + ) { + throw new IllegalArgumentException("datasetId 和 data 不能同时为空"); + } + } + + public static DatasetMeta of(String datasetId) { + return new DatasetMeta(datasetId, "default", "train", null, null); + } + + /** 创建与 Python range 等价的数据切片。 */ + public static JsonObject range(int start, int stop, int step) { + if (step == 0) { + throw new IllegalArgumentException("step 不能为 0"); + } + JsonObject value = new JsonObject(); + value.addProperty("_slice_type_", "range"); + value.addProperty("start", start); + value.addProperty("stop", stop); + value.addProperty("step", step); + return value; + } + + /** 创建由下标列表构成的数据切片。 */ + public static JsonObject indices(List values) { + JsonObject value = new JsonObject(); + value.addProperty("_slice_type_", "list"); + JsonArray array = new JsonArray(); + values.forEach(array::add); + value.add("values", array); + return value; + } + + @Override + public JsonObject toTwinkleJson() { + JsonObject value = new JsonObject(); + value.addProperty("_TWINKLE_TYPE_", "DatasetMeta"); + if (datasetId != null) { + value.addProperty("dataset_id", datasetId); + } + if (subsetName != null) { + value.addProperty("subset_name", subsetName); + } + if (split != null) { + value.addProperty("split", split); + } + if (dataSlice != null) { + value.add("data_slice", new Gson().toJsonTree(dataSlice)); + } + if (data != null) { + value.add("data", new Gson().toJsonTree(data)); + } + return value; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/DeleteCheckpointResult.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/DeleteCheckpointResult.java new file mode 100644 index 00000000..0b523d7e --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/DeleteCheckpointResult.java @@ -0,0 +1,11 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonObject; + +/** 删除检查点后的服务端确认信息。 */ +public record DeleteCheckpointResult( + boolean success, + String message, + JsonObject extensions +) { +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/LoraConfig.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/LoraConfig.java new file mode 100644 index 00000000..359d7ec1 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/LoraConfig.java @@ -0,0 +1,41 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonObject; +import io.github.modelscope.twinkle.transport.TwinkleSerializable; + +/** 与 Python 简化客户端一致的 LoRA 适配器配置。 */ +public record LoraConfig( + int rank, + int loraAlpha, + Object targetModules, + double loraDropout, + String bias, + String taskType +) + implements TwinkleSerializable { + public LoraConfig { + if (rank <= 0) { + throw new IllegalArgumentException("rank 必须大于 0"); + } + } + + /** 使用 Python 客户端相同默认值创建配置。 */ + public LoraConfig() { + this(8, 32, "all-linear", 0.0, "none", null); + } + + @Override + public JsonObject toTwinkleJson() { + JsonObject value = new JsonObject(); + value.addProperty("_TWINKLE_TYPE_", "LoraConfig"); + value.addProperty("r", rank); + value.addProperty("lora_alpha", loraAlpha); + value.add("target_modules", new com.google.gson.Gson().toJsonTree(targetModules)); + value.addProperty("lora_dropout", loraDropout); + value.addProperty("bias", bias); + if (taskType != null) { + value.addProperty("task_type", taskType); + } + return value; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/SampleRequest.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/SampleRequest.java new file mode 100644 index 00000000..5556a1b4 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/SampleRequest.java @@ -0,0 +1,21 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonElement; +import java.util.List; +import java.util.Map; +import java.util.Objects; + +/** 采样请求参数。 */ +public record SampleRequest( + List inputs, + Map samplingParams, + String adapterName, + String adapterUri, + int numSamples +) { + public SampleRequest { + Objects.requireNonNull(inputs, "inputs 不能为空"); + if (numSamples <= 0) throw new IllegalArgumentException("numSamples 必须大于 0"); + adapterName = adapterName == null ? "" : adapterName; + } +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/SampleResult.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/SampleResult.java new file mode 100644 index 00000000..4cee0b6f --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/SampleResult.java @@ -0,0 +1,9 @@ +package io.github.modelscope.twinkle.types; + +import java.util.List; + +/** 单个输入对应的一组采样结果。 */ +public record SampleResult( + List sequences +) { +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/SampledSequence.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/SampledSequence.java new file mode 100644 index 00000000..75dc44fb --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/SampledSequence.java @@ -0,0 +1,12 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonElement; +import java.util.List; + +/** 单条采样序列。 */ +public record SampledSequence( + String stopReason, + List tokens, + String decoded, + JsonElement raw +) {} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/SaveResponse.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/SaveResponse.java new file mode 100644 index 00000000..2bba5b0e --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/SaveResponse.java @@ -0,0 +1,8 @@ +package io.github.modelscope.twinkle.types; + +/** 保存模型或采样器后返回的路径信息。 */ +public record SaveResponse( + String twinklePath, + String checkpointDir +) { +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/ServerCapabilities.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/ServerCapabilities.java new file mode 100644 index 00000000..3927573d --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/ServerCapabilities.java @@ -0,0 +1,11 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonObject; +import java.util.List; + +/** 服务端支持能力的固定描述。 */ +public record ServerCapabilities( + List supportedModels, + JsonObject extensions +) { +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/SupportedModel.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/SupportedModel.java new file mode 100644 index 00000000..14dd7e6b --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/SupportedModel.java @@ -0,0 +1,10 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonObject; + +/** 服务端支持的基础模型。 */ +public record SupportedModel( + String modelName, + JsonObject extensions +) { +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/TrainingRun.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/TrainingRun.java new file mode 100644 index 00000000..520490b8 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/TrainingRun.java @@ -0,0 +1,11 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonObject; + +/** 训练任务摘要;扩展字段保留在 raw 中。 */ +public record TrainingRun( + String trainingRunId, + String baseModel, + String modelOwner, + JsonObject raw +) {} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/TrainingRunPage.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/TrainingRunPage.java new file mode 100644 index 00000000..2eb6583f --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/TrainingRunPage.java @@ -0,0 +1,12 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonObject; +import java.util.List; + +/** 训练任务分页响应。 */ +public record TrainingRunPage( + List runs, + Cursor cursor, + JsonObject extensions +) { +} diff --git a/javaclients/src/main/java/io/github/modelscope/twinkle/types/WeightsInfo.java b/javaclients/src/main/java/io/github/modelscope/twinkle/types/WeightsInfo.java new file mode 100644 index 00000000..d2c05a49 --- /dev/null +++ b/javaclients/src/main/java/io/github/modelscope/twinkle/types/WeightsInfo.java @@ -0,0 +1,13 @@ +package io.github.modelscope.twinkle.types; + +import com.google.gson.JsonObject; + +/** 权重所属训练任务的元数据。 */ +public record WeightsInfo( + String trainingRunId, + String baseModel, + String modelOwner, + boolean isLora, + Integer loraRank, + JsonObject extensions +) {} diff --git a/javaclients/src/test/java/io/github/modelscope/twinkle/LjhTest.java b/javaclients/src/test/java/io/github/modelscope/twinkle/LjhTest.java new file mode 100644 index 00000000..1436c849 --- /dev/null +++ b/javaclients/src/test/java/io/github/modelscope/twinkle/LjhTest.java @@ -0,0 +1,116 @@ +package io.github.modelscope.twinkle; + +import com.google.gson.JsonElement; +import io.github.modelscope.twinkle.types.DatasetKind; +import io.github.modelscope.twinkle.types.DatasetMeta; +import io.github.modelscope.twinkle.types.LoraConfig; +import java.util.Map; + +/** + * 可在 IntelliJ IDEA 中直接运行的完整 LoRA 训练示例。 + * + *

+ * 所有运行参数均来自环境变量,禁止将令牌、内网地址或本机数据路径写入本类。 + *

+ */ +public final class LjhTest { + + private LjhTest() {} + + /** 启动一次端到端的远程 LoRA 训练。 */ + public static void main(String[] args) { + String baseUrl = requiredEnv("TWINKLE_SERVER_URL"); + String token = requiredEnv("TWINKLE_SERVER_TOKEN"); + String modelId = env("TWINKLE_BASE_MODEL", "Qwen/Qwen3.6-27B"); + String datasetId = requiredEnv("TWINKLE_DATASET_ID"); + String template = env("TWINKLE_TEMPLATE", "Qwen3_5Template"); + int batchSize = Integer.parseInt(env("TWINKLE_BATCH_SIZE", "4")); + int epochs = Integer.parseInt(env("TWINKLE_EPOCHS", "1")); + double learningRate = Double.parseDouble(env("TWINKLE_LEARNING_RATE", "0.0001")); + + try ( + TwinkleClient client = TwinkleClient.builder().baseUrl(baseUrl).apiKey(token).build() + ) { + if (!client.healthCheck()) { + throw new IllegalStateException("Twinkle 服务健康检查失败"); + } + + var capabilities = client.serverCapabilities(); + System.out.println("服务支持的模型:"); + capabilities + .supportedModels() + .forEach(supportedModel -> System.out.println("- " + supportedModel.modelName())); + + var capacity = client.capacityInfo(); + System.out.printf( + "LoRA 容量:总数=%d,已用=%d,空闲=%d%n", + capacity.maxLoras(), + capacity.usedLoras(), + capacity.freeLoras() + ); + + var existingRuns = client.trainingRuns().list(10, 0, false); + System.out.printf( + "当前可见训练任务:%d 个%n", + existingRuns.cursor().totalCount() + ); + + var dataset = client + .processors() + .dataset(DatasetKind.DATASET, Map.of("dataset_meta", DatasetMeta.of(datasetId))); + dataset.setTemplate(template, Map.of("model_id", modelId)); + dataset.encode(false, Map.of("batched", true)); + + var dataLoader = client + .processors() + .dataLoader(dataset.processorId(), Map.of("batch_size", batchSize)); + var model = client.models().open(modelId); + model.addAdapter( + "default", + new LoraConfig(8, 16, "all-linear", 0.01, "none", null), + Map.of("gradient_accumulation_steps", 1) + ); + model.setTemplate(template, Map.of()); + model.setProcessor("InputProcessor", Map.of("padding_side", "right")); + model.setLoss("CrossEntropyLoss"); + model.setOptimizer("Adam", Map.of("lr", learningRate)); + + for (int epoch = 0; epoch < epochs; epoch++) { + int step = 0; + for (JsonElement batch : dataLoader) { + model.forwardBackward(batch); + model.clipGradAndStep(1.0, 2); + if (step % 10 == 0) System.out.printf( + "第 %d 轮,第 %d 步,指标:%s%n", + epoch + 1, + step, + model.calculateMetric(true) + ); + step++; + } + System.out.printf("第 %d 轮训练完成,共 %d 步%n", epoch + 1, step); + } + + var saved = model.save("twinkle-java-final", true); + System.out.println("检查点 Twinkle 路径:" + saved.twinklePath()); + if (saved.checkpointDir() != null) { + System.out.println("检查点本地目录:" + saved.checkpointDir()); + } + } + } + + /** 获取必填环境变量。 */ + private static String requiredEnv(String name) { + String value = System.getenv(name); + if (value == null || value.isBlank()) throw new IllegalStateException( + "请设置环境变量 " + name + ); + return value; + } + + /** 获取带默认值的环境变量。 */ + private static String env(String name, String defaultValue) { + String value = System.getenv(name); + return value == null || value.isBlank() ? defaultValue : value; + } +} diff --git a/javaclients/src/test/java/io/github/modelscope/twinkle/ProtocolContractTest.java b/javaclients/src/test/java/io/github/modelscope/twinkle/ProtocolContractTest.java new file mode 100644 index 00000000..57cfd81c --- /dev/null +++ b/javaclients/src/test/java/io/github/modelscope/twinkle/ProtocolContractTest.java @@ -0,0 +1,62 @@ +package io.github.modelscope.twinkle; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import io.github.modelscope.twinkle.types.DatasetMeta; +import io.github.modelscope.twinkle.types.LoraConfig; +import java.util.Map; +import okhttp3.mockwebserver.MockResponse; +import okhttp3.mockwebserver.MockWebServer; +import org.junit.jupiter.api.Test; + +/** 与 Python 简化客户端保持网络协议一致的离线测试。 */ +class ProtocolContractTest { + + /** 管理类接口必须携带 /twinkle 路由前缀。 */ + @Test + void 管理接口应使用Twinkle路由前缀() throws Exception { + try (MockWebServer server = new MockWebServer()) { + server.start(); + try ( + TwinkleClient client = TwinkleClient.builder() + .baseUrl(server.url("/").toString()) + .apiKey("token") + .existingSessionId("session-1") + .build() + ) { + server.enqueue(new MockResponse().setBody("{}")); + client.serverCapabilities(); + assertEquals( + "/api/v1/twinkle/get_server_capabilities", + server.takeRequest().getPath() + ); + } + } + } + + /** DatasetMeta 必须保留 Python 客户端支持的切片与内存数据字段。 */ + @Test + void 数据集元数据应完整序列化() { + var meta = new DatasetMeta( + "", + "default", + "train", + DatasetMeta.range(10, 20, 2), + Map.of("text", "你好") + ); + var json = meta.toTwinkleJson(); + assertEquals("DatasetMeta", json.get("_TWINKLE_TYPE_").getAsString()); + assertEquals("range", json.getAsJsonObject("data_slice").get("_slice_type_").getAsString()); + assertTrue(json.has("data")); + } + + /** LoRA 配置必须使用服务端所需的字段名。 */ + @Test + void LoRA配置应使用Python协议字段名() { + var json = new LoraConfig(16, 32, "all-linear", 0.1, "none", null).toTwinkleJson(); + assertEquals(16, json.get("r").getAsInt()); + assertEquals(32, json.get("lora_alpha").getAsInt()); + assertEquals("all-linear", json.get("target_modules").getAsString()); + } +} diff --git a/javaclients/src/test/java/io/github/modelscope/twinkle/TwinkleClientTest.java b/javaclients/src/test/java/io/github/modelscope/twinkle/TwinkleClientTest.java new file mode 100644 index 00000000..0f466ae0 --- /dev/null +++ b/javaclients/src/test/java/io/github/modelscope/twinkle/TwinkleClientTest.java @@ -0,0 +1,63 @@ +package io.github.modelscope.twinkle; + +import static org.junit.jupiter.api.Assertions.assertNotNull; + +import java.util.Map; +import org.junit.jupiter.api.Assumptions; +import org.junit.jupiter.api.Test; + +/** + * 新版客户端的服务联调测试。 + * + *

仅在同时设置 {@code TWINKLE_SERVER_URL} 和 {@code TWINKLE_SERVER_TOKEN} 时执行; + * 未设置时会自动跳过,避免默认构建访问网络。

+ */ +class TwinkleClientTest { + + /** 验证服务健康检查和会话创建。 */ + @Test + void 应能创建会话并访问健康检查接口() { + try (TwinkleClient client = createClientOrSkip()) { + assertNotNull(client.sessionId()); + assertNotNull(client.healthCheck()); + } + } + + /** 验证服务能力和容量接口可返回 JSON 对象。 */ + @Test + void 应能读取服务能力和容量信息() { + try (TwinkleClient client = createClientOrSkip()) { + assertNotNull(client.serverCapabilities()); + assertNotNull(client.capacityInfo()); + } + } + + /** 验证训练任务列表 API 的新版调用方式。 */ + @Test + void 应能查询训练任务列表() { + try (TwinkleClient client = createClientOrSkip()) { + var page = client.trainingRuns().list(10, 0, false); + assertNotNull(page.runs()); + assertNotNull(page.cursor()); + } + } + + /** 根据环境变量构建客户端;缺少联调配置时跳过测试。 */ + private static TwinkleClient createClientOrSkip() { + String baseUrl = System.getenv("TWINKLE_SERVER_URL"); + String token = System.getenv("TWINKLE_SERVER_TOKEN"); + Assumptions.assumeTrue( + baseUrl != null && !baseUrl.isBlank(), + "未设置 TWINKLE_SERVER_URL,跳过服务联调测试" + ); + Assumptions.assumeTrue( + token != null && !token.isBlank(), + "未设置 TWINKLE_SERVER_TOKEN,跳过服务联调测试" + ); + return TwinkleClient.builder() + .baseUrl(baseUrl) + .apiKey(token) + .sessionMetadata(Map.of("client", "twinkle-client-java-test")) + .build(); + } +} From 4a61b6d4b4285516c68a0d64a3ab99d87240e295 Mon Sep 17 00:00:00 2001 From: liujihui Date: Tue, 4 Aug 2026 14:58:25 +0800 Subject: [PATCH 2/2] =?UTF-8?q?test:=20=E4=BF=AE=E5=A4=8D=E6=9C=8D?= =?UTF-8?q?=E5=8A=A1=E8=83=BD=E5=8A=9B=E5=8D=8F=E8=AE=AE=E5=A5=91=E7=BA=A6?= =?UTF-8?q?=E6=B5=8B=E8=AF=95=E5=93=8D=E5=BA=94?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../java/io/github/modelscope/twinkle/ProtocolContractTest.java | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/javaclients/src/test/java/io/github/modelscope/twinkle/ProtocolContractTest.java b/javaclients/src/test/java/io/github/modelscope/twinkle/ProtocolContractTest.java index 57cfd81c..14d4557f 100644 --- a/javaclients/src/test/java/io/github/modelscope/twinkle/ProtocolContractTest.java +++ b/javaclients/src/test/java/io/github/modelscope/twinkle/ProtocolContractTest.java @@ -25,7 +25,7 @@ class ProtocolContractTest { .existingSessionId("session-1") .build() ) { - server.enqueue(new MockResponse().setBody("{}")); + server.enqueue(new MockResponse().setBody("{\"supported_models\":[]}")); client.serverCapabilities(); assertEquals( "/api/v1/twinkle/get_server_capabilities",