update fmt
This commit is contained in:
Generated
+103
-4
@@ -59,6 +59,7 @@ dependencies = [
|
|||||||
"serde_json",
|
"serde_json",
|
||||||
"serde_yaml",
|
"serde_yaml",
|
||||||
"symphonia",
|
"symphonia",
|
||||||
|
"sysinfo",
|
||||||
"tokenizers",
|
"tokenizers",
|
||||||
"tokio",
|
"tokio",
|
||||||
"url",
|
"url",
|
||||||
@@ -1787,7 +1788,7 @@ dependencies = [
|
|||||||
"libc",
|
"libc",
|
||||||
"log",
|
"log",
|
||||||
"rustversion",
|
"rustversion",
|
||||||
"windows",
|
"windows 0.48.0",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -2119,7 +2120,7 @@ dependencies = [
|
|||||||
"js-sys",
|
"js-sys",
|
||||||
"log",
|
"log",
|
||||||
"wasm-bindgen",
|
"wasm-bindgen",
|
||||||
"windows-core",
|
"windows-core 0.62.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -2855,6 +2856,15 @@ version = "0.3.0"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "0676bb32a98c1a483ce53e500a81ad9c3d5b3f7c920c28c24e9cb0980d0b5bc8"
|
checksum = "0676bb32a98c1a483ce53e500a81ad9c3d5b3f7c920c28c24e9cb0980d0b5bc8"
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "ntapi"
|
||||||
|
version = "0.4.2"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "c70f219e21142367c70c0b30c6a9e3a14d55b4d12a204d897fbec83a0363f081"
|
||||||
|
dependencies = [
|
||||||
|
"winapi",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "nu-ansi-term"
|
name = "nu-ansi-term"
|
||||||
version = "0.50.3"
|
version = "0.50.3"
|
||||||
@@ -4623,6 +4633,20 @@ dependencies = [
|
|||||||
"walkdir",
|
"walkdir",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "sysinfo"
|
||||||
|
version = "0.33.1"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "4fc858248ea01b66f19d8e8a6d55f41deaf91e9d495246fd01368d99935c6c01"
|
||||||
|
dependencies = [
|
||||||
|
"core-foundation-sys",
|
||||||
|
"libc",
|
||||||
|
"memchr",
|
||||||
|
"ntapi",
|
||||||
|
"rayon",
|
||||||
|
"windows 0.57.0",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "system-configuration"
|
name = "system-configuration"
|
||||||
version = "0.6.1"
|
version = "0.6.1"
|
||||||
@@ -5460,6 +5484,22 @@ version = "0.1.10"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "a751b3277700db47d3e574514de2eced5e54dc8a5436a3bf7a0b248b2cee16f3"
|
checksum = "a751b3277700db47d3e574514de2eced5e54dc8a5436a3bf7a0b248b2cee16f3"
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "winapi"
|
||||||
|
version = "0.3.9"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "5c839a674fcd7a98952e593242ea400abe93992746761e38641405d28b00f419"
|
||||||
|
dependencies = [
|
||||||
|
"winapi-i686-pc-windows-gnu",
|
||||||
|
"winapi-x86_64-pc-windows-gnu",
|
||||||
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "winapi-i686-pc-windows-gnu"
|
||||||
|
version = "0.4.0"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "ac3b87c63620426dd9b991e5ce0329eff545bccbbb34f3be09ff6fb6ab51b7b6"
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "winapi-util"
|
name = "winapi-util"
|
||||||
version = "0.1.11"
|
version = "0.1.11"
|
||||||
@@ -5469,6 +5509,12 @@ dependencies = [
|
|||||||
"windows-sys 0.61.2",
|
"windows-sys 0.61.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "winapi-x86_64-pc-windows-gnu"
|
||||||
|
version = "0.4.0"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f"
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "windows"
|
name = "windows"
|
||||||
version = "0.48.0"
|
version = "0.48.0"
|
||||||
@@ -5478,19 +5524,52 @@ dependencies = [
|
|||||||
"windows-targets 0.48.5",
|
"windows-targets 0.48.5",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "windows"
|
||||||
|
version = "0.57.0"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "12342cb4d8e3b046f3d80effd474a7a02447231330ef77d71daa6fbc40681143"
|
||||||
|
dependencies = [
|
||||||
|
"windows-core 0.57.0",
|
||||||
|
"windows-targets 0.52.6",
|
||||||
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "windows-core"
|
||||||
|
version = "0.57.0"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "d2ed2439a290666cd67ecce2b0ffaad89c2a56b976b736e6ece670297897832d"
|
||||||
|
dependencies = [
|
||||||
|
"windows-implement 0.57.0",
|
||||||
|
"windows-interface 0.57.0",
|
||||||
|
"windows-result 0.1.2",
|
||||||
|
"windows-targets 0.52.6",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "windows-core"
|
name = "windows-core"
|
||||||
version = "0.62.2"
|
version = "0.62.2"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "b8e83a14d34d0623b51dce9581199302a221863196a1dde71a7663a4c2be9deb"
|
checksum = "b8e83a14d34d0623b51dce9581199302a221863196a1dde71a7663a4c2be9deb"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"windows-implement",
|
"windows-implement 0.60.2",
|
||||||
"windows-interface",
|
"windows-interface 0.59.3",
|
||||||
"windows-link 0.2.1",
|
"windows-link 0.2.1",
|
||||||
"windows-result 0.4.1",
|
"windows-result 0.4.1",
|
||||||
"windows-strings 0.5.1",
|
"windows-strings 0.5.1",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "windows-implement"
|
||||||
|
version = "0.57.0"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "9107ddc059d5b6fbfbffdfa7a7fe3e22a226def0b2608f72e9d552763d3e1ad7"
|
||||||
|
dependencies = [
|
||||||
|
"proc-macro2",
|
||||||
|
"quote",
|
||||||
|
"syn",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "windows-implement"
|
name = "windows-implement"
|
||||||
version = "0.60.2"
|
version = "0.60.2"
|
||||||
@@ -5502,6 +5581,17 @@ dependencies = [
|
|||||||
"syn",
|
"syn",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "windows-interface"
|
||||||
|
version = "0.57.0"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "29bee4b38ea3cde66011baa44dba677c432a78593e202392d1e9070cf2a7fca7"
|
||||||
|
dependencies = [
|
||||||
|
"proc-macro2",
|
||||||
|
"quote",
|
||||||
|
"syn",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "windows-interface"
|
name = "windows-interface"
|
||||||
version = "0.59.3"
|
version = "0.59.3"
|
||||||
@@ -5536,6 +5626,15 @@ dependencies = [
|
|||||||
"windows-strings 0.4.2",
|
"windows-strings 0.4.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "windows-result"
|
||||||
|
version = "0.1.2"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "5e383302e8ec8515204254685643de10811af0ed97ea37210dc26fb0032647f8"
|
||||||
|
dependencies = [
|
||||||
|
"windows-targets 0.52.6",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "windows-result"
|
name = "windows-result"
|
||||||
version = "0.3.4"
|
version = "0.3.4"
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ hound = "3.5.1"
|
|||||||
clap = { version = "4.5.51", features = ["derive"] }
|
clap = { version = "4.5.51", features = ["derive"] }
|
||||||
modelscope = "0.1.3"
|
modelscope = "0.1.3"
|
||||||
dirs = "6.0.0"
|
dirs = "6.0.0"
|
||||||
|
sysinfo = "0.33"
|
||||||
url = "2.5.7"
|
url = "2.5.7"
|
||||||
rayon = "1.10"
|
rayon = "1.10"
|
||||||
# rubato = "1.0.0"
|
# rubato = "1.0.0"
|
||||||
|
|||||||
+147
@@ -57,6 +57,92 @@ Error responses:
|
|||||||
|
|
||||||
## Endpoints
|
## Endpoints
|
||||||
|
|
||||||
|
### Health Check
|
||||||
|
|
||||||
|
Check the service health status. This endpoint is useful for container orchestration (Kubernetes), load balancers, and monitoring systems.
|
||||||
|
|
||||||
|
#### Endpoint
|
||||||
|
```
|
||||||
|
GET /health
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Response
|
||||||
|
|
||||||
|
**Healthy (HTTP 200):**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"status": "ok"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Unhealthy (HTTP 503):**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"status": "unhealthy",
|
||||||
|
"error": "model not initialized"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Example
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl http://127.0.0.1:10100/health
|
||||||
|
```
|
||||||
|
|
||||||
|
### Models
|
||||||
|
|
||||||
|
Get information about the currently loaded model (OpenAI API compatible format).
|
||||||
|
|
||||||
|
#### Endpoint
|
||||||
|
```
|
||||||
|
GET /models
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Response
|
||||||
|
|
||||||
|
**Success (HTTP 200):**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"object": "list",
|
||||||
|
"data": [
|
||||||
|
{
|
||||||
|
"id": "qwen3-0.6b",
|
||||||
|
"object": "model",
|
||||||
|
"created": null,
|
||||||
|
"owned_by": "Qwen"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Not Initialized (HTTP 503):**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"error": "model not initialized"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Fields
|
||||||
|
|
||||||
|
| Field | Type | Description |
|
||||||
|
|-------|------|-------------|
|
||||||
|
| `object` | string | Fixed value: "list" |
|
||||||
|
| `data` | array | Array of model objects (currently contains one loaded model) |
|
||||||
|
| `id` | string | Model identifier in kebab-case (e.g., "qwen3-0.6b") |
|
||||||
|
| `object` | string | Fixed value: "model" |
|
||||||
|
| `created` | integer\|null | Unix timestamp (currently null) |
|
||||||
|
| `owned_by` | string | Model owner/organization name |
|
||||||
|
|
||||||
|
#### Example
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl http://127.0.0.1:10100/models
|
||||||
|
```
|
||||||
|
|
||||||
### Chat Completions
|
### Chat Completions
|
||||||
|
|
||||||
Generate chat completions or text responses.
|
Generate chat completions or text responses.
|
||||||
@@ -352,6 +438,67 @@ Returns the processed image in base64 PNG format.
|
|||||||
|
|
||||||
- `rmbg2.0`
|
- `rmbg2.0`
|
||||||
|
|
||||||
|
### Graceful Shutdown
|
||||||
|
|
||||||
|
Gracefully shut down the AHA server. This endpoint initiates a graceful shutdown process that:
|
||||||
|
1. Stops accepting new connections
|
||||||
|
2. Waits for existing requests to complete (up to 1 second)
|
||||||
|
3. Cleans up PID files
|
||||||
|
4. Exits the process
|
||||||
|
|
||||||
|
#### Endpoint
|
||||||
|
```
|
||||||
|
POST /shutdown
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Request Body
|
||||||
|
|
||||||
|
None (empty request)
|
||||||
|
|
||||||
|
#### Response
|
||||||
|
|
||||||
|
**Success (HTTP 200):**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"message": "Shutting down..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Forbidden (HTTP 403):**
|
||||||
|
|
||||||
|
When remote shutdown is not allowed:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"error": "Remote shutdown not allowed. Use --allow-remote-shutdown flag to enable (not recommended)."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Security
|
||||||
|
|
||||||
|
By default, the shutdown endpoint only allows requests from localhost (127.0.0.1). To enable remote shutdown, start the server with the `--allow-remote-shutdown` flag:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
aha serv -m qwen3-0.6b --allow-remote-shutdown
|
||||||
|
```
|
||||||
|
|
||||||
|
**Warning:** Enabling remote shutdown is not recommended for production use unless properly secured.
|
||||||
|
|
||||||
|
#### Example
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://127.0.0.1:10100/shutdown
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Logging
|
||||||
|
|
||||||
|
All shutdown requests are logged to stderr with the format:
|
||||||
|
|
||||||
|
```
|
||||||
|
[SHUTDOWN] Shutdown requested (remote_allowed: false)
|
||||||
|
```
|
||||||
|
|
||||||
## Error Handling
|
## Error Handling
|
||||||
|
|
||||||
### Error Codes
|
### Error Codes
|
||||||
|
|||||||
@@ -57,6 +57,92 @@ Content-Type: application/json
|
|||||||
|
|
||||||
## 端点
|
## 端点
|
||||||
|
|
||||||
|
### 健康检查
|
||||||
|
|
||||||
|
检查服务健康状态。此端点适用于容器编排(Kubernetes)、负载均衡器和监控系统。
|
||||||
|
|
||||||
|
#### 端点
|
||||||
|
```
|
||||||
|
GET /health
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 响应
|
||||||
|
|
||||||
|
**健康 (HTTP 200):**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"status": "ok"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**不健康 (HTTP 503):**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"status": "unhealthy",
|
||||||
|
"error": "model not initialized"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 示例
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl http://127.0.0.1:10100/health
|
||||||
|
```
|
||||||
|
|
||||||
|
### 模型列表
|
||||||
|
|
||||||
|
获取当前加载的模型信息(OpenAI API 兼容格式)。
|
||||||
|
|
||||||
|
#### 端点
|
||||||
|
```
|
||||||
|
GET /models
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 响应
|
||||||
|
|
||||||
|
**成功 (HTTP 200):**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"object": "list",
|
||||||
|
"data": [
|
||||||
|
{
|
||||||
|
"id": "qwen3-0.6b",
|
||||||
|
"object": "model",
|
||||||
|
"created": null,
|
||||||
|
"owned_by": "Qwen"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**未初始化 (HTTP 503):**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"error": "model not initialized"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 字段
|
||||||
|
|
||||||
|
| 字段 | 类型 | 描述 |
|
||||||
|
|------|------|------|
|
||||||
|
| `object` | string | 固定值:"list" |
|
||||||
|
| `data` | array | 模型对象数组(当前仅包含一个已加载的模型) |
|
||||||
|
| `id` | string | 模型标识符(kebab-case,如 "qwen3-0.6b") |
|
||||||
|
| `object` | string | 固定值:"model" |
|
||||||
|
| `created` | integer\|null | Unix 时间戳(当前为 null) |
|
||||||
|
| `owned_by` | string | 模型所有者/组织名称 |
|
||||||
|
|
||||||
|
#### 示例
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl http://127.0.0.1:10100/models
|
||||||
|
```
|
||||||
|
|
||||||
### 对话补全
|
### 对话补全
|
||||||
|
|
||||||
生成对话补全或文本响应。
|
生成对话补全或文本响应。
|
||||||
@@ -352,6 +438,67 @@ curl http://127.0.0.1:10100/images/remove_background \
|
|||||||
|
|
||||||
- `rmbg2.0`
|
- `rmbg2.0`
|
||||||
|
|
||||||
|
### 优雅关机
|
||||||
|
|
||||||
|
优雅地关闭 AHA 服务器。此端点启动优雅关闭流程:
|
||||||
|
1. 停止接受新连接
|
||||||
|
2. 等待现有请求完成(最多 1 秒)
|
||||||
|
3. 清理 PID 文件
|
||||||
|
4. 退出进程
|
||||||
|
|
||||||
|
#### 端点
|
||||||
|
```
|
||||||
|
POST /shutdown
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 请求体
|
||||||
|
|
||||||
|
无(空请求)
|
||||||
|
|
||||||
|
#### 响应
|
||||||
|
|
||||||
|
**成功 (HTTP 200):**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"message": "Shutting down..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**禁止访问 (HTTP 403):**
|
||||||
|
|
||||||
|
当不允许远程关闭时:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"error": "Remote shutdown not allowed. Use --allow-remote-shutdown flag to enable (not recommended)."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 安全性
|
||||||
|
|
||||||
|
默认情况下,关机端点仅允许来自 localhost (127.0.0.1) 的请求。要启用远程关闭,请使用 `--allow-remote-shutdown` 标志启动服务器:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
aha serv -m qwen3-0.6b --allow-remote-shutdown
|
||||||
|
```
|
||||||
|
|
||||||
|
**警告:** 除非有适当的安全措施,否则不建议在生产环境中启用远程关闭。
|
||||||
|
|
||||||
|
#### 示例
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://127.0.0.1:10100/shutdown
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 日志记录
|
||||||
|
|
||||||
|
所有关机请求都会记录到 stderr,格式如下:
|
||||||
|
|
||||||
|
```
|
||||||
|
[SHUTDOWN] Shutdown requested (remote_allowed: false)
|
||||||
|
```
|
||||||
|
|
||||||
## 错误处理
|
## 错误处理
|
||||||
|
|
||||||
### 错误代码
|
### 错误代码
|
||||||
|
|||||||
+148
-5
@@ -113,11 +113,11 @@ aha run -m qwen3asr-0.6b -i "audio.wav" --weight-path /path/to/model
|
|||||||
|
|
||||||
### serv - Start service
|
### serv - Start service
|
||||||
|
|
||||||
Start HTTP service only, without downloading models. Must specify local model path via `--weight-path`.
|
Start HTTP service with a model. The `--weight-path` is optional - if not specified, it defaults to `~/.aha/{model_id}`.
|
||||||
|
|
||||||
**Syntax:**
|
**Syntax:**
|
||||||
```bash
|
```bash
|
||||||
aha serv [OPTIONS] --model <MODEL> --weight-path <WEIGHT_PATH>
|
aha serv [OPTIONS] --model <MODEL> [--weight-path <WEIGHT_PATH>]
|
||||||
```
|
```
|
||||||
|
|
||||||
**Options:**
|
**Options:**
|
||||||
@@ -127,21 +127,69 @@ aha serv [OPTIONS] --model <MODEL> --weight-path <WEIGHT_PATH>
|
|||||||
| `-a, --address <ADDRESS>` | Service listen address | 127.0.0.1 |
|
| `-a, --address <ADDRESS>` | Service listen address | 127.0.0.1 |
|
||||||
| `-p, --port <PORT>` | Service listen port | 10100 |
|
| `-p, --port <PORT>` | Service listen port | 10100 |
|
||||||
| `-m, --model <MODEL>` | Model type (required) | - |
|
| `-m, --model <MODEL>` | Model type (required) | - |
|
||||||
| `--weight-path <WEIGHT_PATH>` | Local model weight path (required) | - |
|
| `--weight-path <WEIGHT_PATH>` | Local model weight path (optional) | ~/.aha/{model_id} |
|
||||||
|
| `--allow-remote-shutdown` | Allow remote shutdown requests (not recommended) | false |
|
||||||
|
|
||||||
**Examples:**
|
**Examples:**
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
# Start service with default model path (~/.aha/{model_id})
|
||||||
|
aha serv -m qwen3vl-2b
|
||||||
|
|
||||||
# Start service with local model
|
# Start service with local model
|
||||||
aha serv -m qwen3vl-2b --weight-path /path/to/model
|
aha serv -m qwen3vl-2b --weight-path /path/to/model
|
||||||
|
|
||||||
# Start with specified port
|
# Start with specified port
|
||||||
aha serv -m qwen3vl-2b --weight-path /path/to/model -p 8080
|
aha serv -m qwen3vl-2b -p 8080
|
||||||
|
|
||||||
# Specify listen address
|
# Specify listen address
|
||||||
aha serv -m qwen3vl-2b --weight-path /path/to/model -a 0.0.0.0
|
aha serv -m qwen3vl-2b -a 0.0.0.0
|
||||||
|
|
||||||
|
# Enable remote shutdown (not recommended for production)
|
||||||
|
aha serv -m qwen3vl-2b --allow-remote-shutdown
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### ps - List running services
|
||||||
|
|
||||||
|
List all currently running AHA services with their process IDs, ports, and status.
|
||||||
|
|
||||||
|
**Syntax:**
|
||||||
|
```bash
|
||||||
|
aha ps [OPTIONS]
|
||||||
|
```
|
||||||
|
|
||||||
|
**Options:**
|
||||||
|
|
||||||
|
| Option | Description | Default |
|
||||||
|
|--------|-------------|---------|
|
||||||
|
| `-c, --compact` | Compact output format (show service IDs only) | false |
|
||||||
|
|
||||||
|
**Examples:**
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# List all running services (table format)
|
||||||
|
aha ps
|
||||||
|
|
||||||
|
# Compact output (service IDs only)
|
||||||
|
aha ps -c
|
||||||
|
```
|
||||||
|
|
||||||
|
**Output Format:**
|
||||||
|
|
||||||
|
```
|
||||||
|
Service ID PID Model Port Address Status
|
||||||
|
-------------------------------------------------------------------------------------
|
||||||
|
56860@10100 56860 N/A 10100 127.0.0.1 Running
|
||||||
|
```
|
||||||
|
|
||||||
|
**Fields:**
|
||||||
|
- `Service ID`: Unique identifier in format `pid@port`
|
||||||
|
- `PID`: Process ID
|
||||||
|
- `Model`: Model name (N/A if not detected)
|
||||||
|
- `Port`: Service port number
|
||||||
|
- `Address`: Service listen address
|
||||||
|
- `Status`: Service status (Running, Stopping, Unknown)
|
||||||
|
|
||||||
### download - Download model
|
### download - Download model
|
||||||
|
|
||||||
Download the specified model only, without starting the service.
|
Download the specified model only, without starting the service.
|
||||||
@@ -175,6 +223,94 @@ aha download -m qwen3vl-2b --download-retries 5
|
|||||||
aha download -m minicpm4-0.5b -s models
|
aha download -m minicpm4-0.5b -s models
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### delete - Delete downloaded model
|
||||||
|
|
||||||
|
Delete a downloaded model from the default location (`~/.aha/{model_id}`).
|
||||||
|
|
||||||
|
**Syntax:**
|
||||||
|
```bash
|
||||||
|
aha delete [OPTIONS] --model <MODEL>
|
||||||
|
```
|
||||||
|
|
||||||
|
**Options:**
|
||||||
|
|
||||||
|
| Option | Description | Default |
|
||||||
|
|--------|-------------|---------|
|
||||||
|
| `-m, --model <MODEL>` | Model type (required) | - |
|
||||||
|
|
||||||
|
**Examples:**
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Delete RMBG2.0 model from default location
|
||||||
|
aha delete -m rmbg2.0
|
||||||
|
|
||||||
|
# Delete Qwen3-VL-2B model
|
||||||
|
aha delete --model qwen3vl-2b
|
||||||
|
```
|
||||||
|
|
||||||
|
**Behavior:**
|
||||||
|
- Displays model information (ID, location, size) before deletion
|
||||||
|
- Requires confirmation (y/N) before proceeding
|
||||||
|
- Shows "Model not found" message if the model directory doesn't exist
|
||||||
|
- Shows "Model deleted successfully" message after completion
|
||||||
|
|
||||||
|
### list - List all supported models
|
||||||
|
|
||||||
|
List all supported models with their ModelScope IDs.
|
||||||
|
|
||||||
|
**Syntax:**
|
||||||
|
```bash
|
||||||
|
aha list [OPTIONS]
|
||||||
|
```
|
||||||
|
|
||||||
|
**Options:**
|
||||||
|
|
||||||
|
| Option | Description | Default |
|
||||||
|
|--------|-------------|---------|
|
||||||
|
| `-j, --json` | Output in JSON format (includes name, model_id, and type fields) | false |
|
||||||
|
|
||||||
|
**Examples:**
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# List models in table format (default)
|
||||||
|
aha list
|
||||||
|
|
||||||
|
# List models in JSON format
|
||||||
|
aha list --json
|
||||||
|
|
||||||
|
# Short form
|
||||||
|
aha list -j
|
||||||
|
```
|
||||||
|
|
||||||
|
**JSON Output Format:**
|
||||||
|
|
||||||
|
When using `--json`, the output includes:
|
||||||
|
- `name`: Model identifier used with `-m` flag
|
||||||
|
- `model_id`: Full ModelScope model ID
|
||||||
|
- `type`: Model category (`llm`, `ocr`, `asr`, or `image`)
|
||||||
|
|
||||||
|
Example:
|
||||||
|
```json
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"name": "qwen3vl-2b",
|
||||||
|
"model_id": "Qwen/Qwen3-VL-2B-Instruct",
|
||||||
|
"type": "llm"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "deepseek-ocr",
|
||||||
|
"model_id": "deepseek-ai/DeepSeek-OCR",
|
||||||
|
"type": "ocr"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
```
|
||||||
|
|
||||||
|
**Model Types:**
|
||||||
|
- `llm`: Language models (text generation, chat, etc.)
|
||||||
|
- `ocr`: Optical Character Recognition models
|
||||||
|
- `asr`: Automatic Speech Recognition models
|
||||||
|
- `image`: Image processing models
|
||||||
|
|
||||||
## Supported Models
|
## Supported Models
|
||||||
|
|
||||||
| Model ID | Model Name | Description |
|
| Model ID | Model Name | Description |
|
||||||
@@ -254,6 +390,13 @@ After the service starts, the following API endpoints are available:
|
|||||||
- **Format**: OpenAI Chat Completion format
|
- **Format**: OpenAI Chat Completion format
|
||||||
- **Streaming Support**: No
|
- **Streaming Support**: No
|
||||||
|
|
||||||
|
### Shutdown Endpoint
|
||||||
|
- **Endpoint**: `POST /shutdown`
|
||||||
|
- **Function**: Gracefully shut down the server
|
||||||
|
- **Security**: Localhost only by default, use `--allow-remote-shutdown` flag to enable remote access (not recommended)
|
||||||
|
- **Format**: JSON response
|
||||||
|
|
||||||
|
|
||||||
## Backward Compatibility
|
## Backward Compatibility
|
||||||
|
|
||||||
To maintain compatibility with older versions, the following two usage methods are equivalent:
|
To maintain compatibility with older versions, the following two usage methods are equivalent:
|
||||||
|
|||||||
+147
-5
@@ -113,11 +113,11 @@ aha run -m qwen3asr-0.6b -i "audio.wav" --weight-path /path/to/model
|
|||||||
|
|
||||||
### serv - 启动服务
|
### serv - 启动服务
|
||||||
|
|
||||||
仅启动 HTTP 服务,不下载模型。必须通过 `--weight-path` 指定本地模型路径。
|
使用指定模型启动 HTTP 服务。`--weight-path` 是可选的 - 如果不指定,默认使用 `~/.aha/{model_id}`。
|
||||||
|
|
||||||
**语法:**
|
**语法:**
|
||||||
```bash
|
```bash
|
||||||
aha serv [OPTIONS] --model <MODEL> --weight-path <WEIGHT_PATH>
|
aha serv [OPTIONS] --model <MODEL> [--weight-path <WEIGHT_PATH>]
|
||||||
```
|
```
|
||||||
|
|
||||||
**选项:**
|
**选项:**
|
||||||
@@ -127,21 +127,69 @@ aha serv [OPTIONS] --model <MODEL> --weight-path <WEIGHT_PATH>
|
|||||||
| `-a, --address <ADDRESS>` | 服务监听地址 | 127.0.0.1 |
|
| `-a, --address <ADDRESS>` | 服务监听地址 | 127.0.0.1 |
|
||||||
| `-p, --port <PORT>` | 服务监听端口 | 10100 |
|
| `-p, --port <PORT>` | 服务监听端口 | 10100 |
|
||||||
| `-m, --model <MODEL>` | 模型类型(必选) | - |
|
| `-m, --model <MODEL>` | 模型类型(必选) | - |
|
||||||
| `--weight-path <WEIGHT_PATH>` | 本地模型权重路径(必选) | - |
|
| `--weight-path <WEIGHT_PATH>` | 本地模型权重路径(可选) | ~/.aha/{model_id} |
|
||||||
|
| `--allow-remote-shutdown` | 允许远程关机请求(不推荐) | false |
|
||||||
|
|
||||||
**示例:**
|
**示例:**
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
# 使用默认模型路径启动服务 (~/.aha/{model_id})
|
||||||
|
aha serv -m qwen3vl-2b
|
||||||
|
|
||||||
# 使用本地模型启动服务
|
# 使用本地模型启动服务
|
||||||
aha serv -m qwen3vl-2b --weight-path /path/to/model
|
aha serv -m qwen3vl-2b --weight-path /path/to/model
|
||||||
|
|
||||||
# 指定端口启动
|
# 指定端口启动
|
||||||
aha serv -m qwen3vl-2b --weight-path /path/to/model -p 8080
|
aha serv -m qwen3vl-2b -p 8080
|
||||||
|
|
||||||
# 指定监听地址
|
# 指定监听地址
|
||||||
aha serv -m qwen3vl-2b --weight-path /path/to/model -a 0.0.0.0
|
aha serv -m qwen3vl-2b -a 0.0.0.0
|
||||||
|
|
||||||
|
# 启用远程关机(不推荐用于生产环境)
|
||||||
|
aha serv -m qwen3vl-2b --allow-remote-shutdown
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### ps - 列出运行中的服务
|
||||||
|
|
||||||
|
列出所有当前正在运行的 AHA 服务,显示进程 ID、端口和状态。
|
||||||
|
|
||||||
|
**语法:**
|
||||||
|
```bash
|
||||||
|
aha ps [OPTIONS]
|
||||||
|
```
|
||||||
|
|
||||||
|
**选项:**
|
||||||
|
|
||||||
|
| 选项 | 说明 | 默认值 |
|
||||||
|
|------|------|--------|
|
||||||
|
| `-c, --compact` | 紧凑输出格式(仅显示服务 ID) | false |
|
||||||
|
|
||||||
|
**示例:**
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 列出所有运行中的服务(表格格式)
|
||||||
|
aha ps
|
||||||
|
|
||||||
|
# 紧凑输出(仅服务 ID)
|
||||||
|
aha ps -c
|
||||||
|
```
|
||||||
|
|
||||||
|
**输出格式:**
|
||||||
|
|
||||||
|
```
|
||||||
|
Service ID PID Model Port Address Status
|
||||||
|
-------------------------------------------------------------------------------------
|
||||||
|
56860@10100 56860 N/A 10100 127.0.0.1 Running
|
||||||
|
```
|
||||||
|
|
||||||
|
**字段说明:**
|
||||||
|
- `Service ID`: 服务唯一标识符,格式为 `pid@port`
|
||||||
|
- `PID`: 进程 ID
|
||||||
|
- `Model`: 模型名称(如果未检测到则显示 N/A)
|
||||||
|
- `Port`: 服务端口号
|
||||||
|
- `Address`: 服务监听地址
|
||||||
|
- `Status`: 服务状态(Running、Stopping、Unknown)
|
||||||
|
|
||||||
### download - 下载模型
|
### download - 下载模型
|
||||||
|
|
||||||
仅下载指定模型,不启动服务。
|
仅下载指定模型,不启动服务。
|
||||||
@@ -175,6 +223,94 @@ aha download -m qwen3vl-2b --download-retries 5
|
|||||||
aha download -m minicpm4-0.5b -s models
|
aha download -m minicpm4-0.5b -s models
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### delete - 删除已下载的模型
|
||||||
|
|
||||||
|
删除默认位置(`~/.aha/{model_id}`)的已下载模型。
|
||||||
|
|
||||||
|
**语法:**
|
||||||
|
```bash
|
||||||
|
aha delete [OPTIONS] --model <MODEL>
|
||||||
|
```
|
||||||
|
|
||||||
|
**选项:**
|
||||||
|
|
||||||
|
| 选项 | 说明 | 默认值 |
|
||||||
|
|------|------|--------|
|
||||||
|
| `-m, --model <MODEL>` | 模型类型(必选) | - |
|
||||||
|
|
||||||
|
**示例:**
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 删除 RMBG2.0 模型
|
||||||
|
aha delete -m rmbg2.0
|
||||||
|
|
||||||
|
# 删除 Qwen3-VL-2B 模型
|
||||||
|
aha delete --model qwen3vl-2b
|
||||||
|
```
|
||||||
|
|
||||||
|
**行为说明:**
|
||||||
|
- 删除前会显示模型信息(ID、位置、大小)
|
||||||
|
- 需要用户确认(y/N)才会执行删除
|
||||||
|
- 如果模型目录不存在,显示"模型未找到"消息
|
||||||
|
- 删除完成后显示"删除成功"消息
|
||||||
|
|
||||||
|
### list - 列出所有支持的模型
|
||||||
|
|
||||||
|
列出所有支持的模型及其 ModelScope ID。
|
||||||
|
|
||||||
|
**语法:**
|
||||||
|
```bash
|
||||||
|
aha list [OPTIONS]
|
||||||
|
```
|
||||||
|
|
||||||
|
**选项:**
|
||||||
|
|
||||||
|
| 选项 | 说明 | 默认值 |
|
||||||
|
|------|------|--------|
|
||||||
|
| `-j, --json` | 以 JSON 格式输出(包含 name、model_id 和 type 字段) | false |
|
||||||
|
|
||||||
|
**示例:**
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 以表格格式列出模型(默认)
|
||||||
|
aha list
|
||||||
|
|
||||||
|
# 以 JSON 格式列出模型
|
||||||
|
aha list --json
|
||||||
|
|
||||||
|
# 简写形式
|
||||||
|
aha list -j
|
||||||
|
```
|
||||||
|
|
||||||
|
**JSON 输出格式:**
|
||||||
|
|
||||||
|
使用 `--json` 时,输出包含:
|
||||||
|
- `name`:与 `-m` 参数一起使用的模型标识符
|
||||||
|
- `model_id`:完整的 ModelScope 模型 ID
|
||||||
|
- `type`:模型类别(`llm`、`ocr`、`asr` 或 `image`)
|
||||||
|
|
||||||
|
示例:
|
||||||
|
```json
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"name": "qwen3vl-2b",
|
||||||
|
"model_id": "Qwen/Qwen3-VL-2B-Instruct",
|
||||||
|
"type": "llm"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "deepseek-ocr",
|
||||||
|
"model_id": "deepseek-ai/DeepSeek-OCR",
|
||||||
|
"type": "ocr"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
```
|
||||||
|
|
||||||
|
**模型类型:**
|
||||||
|
- `llm`:语言模型(文本生成、对话等)
|
||||||
|
- `ocr`:光学字符识别模型
|
||||||
|
- `asr`:自动语音识别模型
|
||||||
|
- `image`:图像处理模型
|
||||||
|
|
||||||
## 支持的模型
|
## 支持的模型
|
||||||
|
|
||||||
| 模型标识 | 模型名称 | 说明 |
|
| 模型标识 | 模型名称 | 说明 |
|
||||||
@@ -254,6 +390,12 @@ aha -m qwen3vl-2b -a 0.0.0.0 -p 8080
|
|||||||
- **格式**: OpenAI Chat Completion 格式
|
- **格式**: OpenAI Chat Completion 格式
|
||||||
- **流式支持**: 不支持
|
- **流式支持**: 不支持
|
||||||
|
|
||||||
|
### 关机接口
|
||||||
|
- **端点**: `POST /shutdown`
|
||||||
|
- **功能**: 优雅地关闭服务器
|
||||||
|
- **安全性**: 默认仅允许本地访问,使用 `--allow-remote-shutdown` 标志启用远程访问(不推荐)
|
||||||
|
- **格式**: JSON 响应
|
||||||
|
|
||||||
## 向后兼容性
|
## 向后兼容性
|
||||||
|
|
||||||
为了保持与旧版本的兼容性,以下两种使用方式是等效的:
|
为了保持与旧版本的兼容性,以下两种使用方式是等效的:
|
||||||
|
|||||||
+330
-8
@@ -1,29 +1,59 @@
|
|||||||
use std::pin::pin;
|
use std::pin::pin;
|
||||||
|
use std::sync::atomic::{AtomicBool, Ordering};
|
||||||
use std::sync::{Arc, OnceLock};
|
use std::sync::{Arc, OnceLock};
|
||||||
|
|
||||||
use aha::models::{GenerateModel, ModelInstance, WhichModel, load_model};
|
use aha::models::{GenerateModel, ModelInstance, WhichModel, load_model};
|
||||||
|
use aha::process::cleanup_pid_file;
|
||||||
use aha::utils::string_to_static_str;
|
use aha::utils::string_to_static_str;
|
||||||
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||||
use rocket::futures::StreamExt;
|
use rocket::futures::StreamExt;
|
||||||
use rocket::serde::json::Json;
|
use rocket::serde::{Serialize, json::Json};
|
||||||
use rocket::{
|
use rocket::{
|
||||||
Request,
|
Request, State,
|
||||||
futures::Stream,
|
futures::Stream,
|
||||||
|
get,
|
||||||
http::{ContentType, Status},
|
http::{ContentType, Status},
|
||||||
post,
|
post,
|
||||||
response::{Responder, stream::TextStream},
|
response::{Responder, stream::TextStream},
|
||||||
};
|
};
|
||||||
use tokio::sync::RwLock;
|
use tokio::sync::RwLock;
|
||||||
|
|
||||||
static MODEL: OnceLock<Arc<RwLock<ModelInstance<'static>>>> = OnceLock::new();
|
/// Wrapper to store model type together with the model instance
|
||||||
|
struct StoredModel {
|
||||||
|
which_model: WhichModel,
|
||||||
|
instance: ModelInstance<'static>,
|
||||||
|
}
|
||||||
|
|
||||||
|
static MODEL: OnceLock<Arc<RwLock<StoredModel>>> = OnceLock::new();
|
||||||
|
static SHUTDOWN_FLAG: OnceLock<Arc<AtomicBool>> = OnceLock::new();
|
||||||
|
static SERVER_PORT: OnceLock<u16> = OnceLock::new();
|
||||||
|
static ALLOW_REMOTE_SHUTDOWN: OnceLock<bool> = OnceLock::new();
|
||||||
|
|
||||||
pub fn init(model_type: WhichModel, path: String) -> anyhow::Result<()> {
|
pub fn init(model_type: WhichModel, path: String) -> anyhow::Result<()> {
|
||||||
let model_path = string_to_static_str(path);
|
let model_path = string_to_static_str(path);
|
||||||
let model = load_model(model_type, model_path)?;
|
let model = load_model(model_type, model_path)?;
|
||||||
MODEL.get_or_init(|| Arc::new(RwLock::new(model)));
|
MODEL.get_or_init(|| {
|
||||||
|
Arc::new(RwLock::new(StoredModel {
|
||||||
|
which_model: model_type,
|
||||||
|
instance: model,
|
||||||
|
}))
|
||||||
|
});
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn set_server_port(port: u16, allow_remote_shutdown: bool) {
|
||||||
|
SHUTDOWN_FLAG.get_or_init(|| Arc::new(AtomicBool::new(false)));
|
||||||
|
SERVER_PORT.get_or_init(|| port);
|
||||||
|
ALLOW_REMOTE_SHUTDOWN.get_or_init(|| allow_remote_shutdown);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[allow(unused)]
|
||||||
|
pub fn get_shutdown_flag() -> Arc<AtomicBool> {
|
||||||
|
SHUTDOWN_FLAG
|
||||||
|
.get_or_init(|| Arc::new(AtomicBool::new(false)))
|
||||||
|
.clone()
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) enum Response<R: Stream<Item = String> + Send> {
|
pub(crate) enum Response<R: Stream<Item = String> + Send> {
|
||||||
Stream(TextStream<R>),
|
Stream(TextStream<R>),
|
||||||
Text(String),
|
Text(String),
|
||||||
@@ -62,7 +92,8 @@ pub(crate) async fn chat(
|
|||||||
.cloned()
|
.cloned()
|
||||||
.ok_or_else(|| anyhow::anyhow!("model not init"))
|
.ok_or_else(|| anyhow::anyhow!("model not init"))
|
||||||
.unwrap();
|
.unwrap();
|
||||||
model_ref.write().await.generate(req.into_inner())
|
let mut guard = model_ref.write().await;
|
||||||
|
guard.instance.generate(req.into_inner())
|
||||||
};
|
};
|
||||||
match response {
|
match response {
|
||||||
Ok(res) => {
|
Ok(res) => {
|
||||||
@@ -76,7 +107,7 @@ pub(crate) async fn chat(
|
|||||||
let text_stream = TextStream! {
|
let text_stream = TextStream! {
|
||||||
let model_ref = MODEL.get().cloned().ok_or_else(|| anyhow::anyhow!("model not init")).unwrap();
|
let model_ref = MODEL.get().cloned().ok_or_else(|| anyhow::anyhow!("model not init")).unwrap();
|
||||||
let mut guard = model_ref.write().await;
|
let mut guard = model_ref.write().await;
|
||||||
let stream_result = guard.generate_stream(req.into_inner());
|
let stream_result = guard.instance.generate_stream(req.into_inner());
|
||||||
match stream_result {
|
match stream_result {
|
||||||
Ok(stream) => {
|
Ok(stream) => {
|
||||||
let mut stream = pin!(stream);
|
let mut stream = pin!(stream);
|
||||||
@@ -113,7 +144,8 @@ pub(crate) async fn remove_background(req: Json<ChatCompletionParameters>) -> (S
|
|||||||
.cloned()
|
.cloned()
|
||||||
.ok_or_else(|| anyhow::anyhow!("model not init"))
|
.ok_or_else(|| anyhow::anyhow!("model not init"))
|
||||||
.unwrap();
|
.unwrap();
|
||||||
model_ref.write().await.generate(req.into_inner())
|
let mut guard = model_ref.write().await;
|
||||||
|
guard.instance.generate(req.into_inner())
|
||||||
};
|
};
|
||||||
match response {
|
match response {
|
||||||
Ok(res) => {
|
Ok(res) => {
|
||||||
@@ -132,7 +164,8 @@ pub(crate) async fn speech(req: Json<ChatCompletionParameters>) -> (Status, Stri
|
|||||||
.cloned()
|
.cloned()
|
||||||
.ok_or_else(|| anyhow::anyhow!("model not init"))
|
.ok_or_else(|| anyhow::anyhow!("model not init"))
|
||||||
.unwrap();
|
.unwrap();
|
||||||
model_ref.write().await.generate(req.into_inner())
|
let mut guard = model_ref.write().await;
|
||||||
|
guard.instance.generate(req.into_inner())
|
||||||
};
|
};
|
||||||
match response {
|
match response {
|
||||||
Ok(res) => {
|
Ok(res) => {
|
||||||
@@ -142,3 +175,292 @@ pub(crate) async fn speech(req: Json<ChatCompletionParameters>) -> (Status, Stri
|
|||||||
Err(e) => (Status::InternalServerError, e.to_string()),
|
Err(e) => (Status::InternalServerError, e.to_string()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Health check endpoint
|
||||||
|
|
||||||
|
#[derive(Serialize)]
|
||||||
|
pub(crate) struct HealthResponse {
|
||||||
|
status: String,
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
error: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[get("/health")]
|
||||||
|
pub(crate) async fn health() -> (Status, (ContentType, Json<HealthResponse>)) {
|
||||||
|
if MODEL.get().is_some() {
|
||||||
|
let response = HealthResponse {
|
||||||
|
status: "ok".to_string(),
|
||||||
|
error: None,
|
||||||
|
};
|
||||||
|
(Status::Ok, (ContentType::JSON, Json(response)))
|
||||||
|
} else {
|
||||||
|
let response = HealthResponse {
|
||||||
|
status: "unhealthy".to_string(),
|
||||||
|
error: Some("model not initialized".to_string()),
|
||||||
|
};
|
||||||
|
(
|
||||||
|
Status::ServiceUnavailable,
|
||||||
|
(ContentType::JSON, Json(response)),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Models endpoint (OpenAI-compatible format)
|
||||||
|
|
||||||
|
/// OpenAI-compatible model object
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct ModelObject {
|
||||||
|
id: String,
|
||||||
|
object: String,
|
||||||
|
created: Option<i64>,
|
||||||
|
owned_by: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// OpenAI-compatible models list response
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct ModelsListResponse {
|
||||||
|
object: String,
|
||||||
|
data: Vec<ModelObject>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct ErrorResponse {
|
||||||
|
error: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Convert WhichModel to a display-friendly model ID (kebab-case)
|
||||||
|
fn which_model_to_id(which_model: WhichModel) -> &'static str {
|
||||||
|
match which_model {
|
||||||
|
WhichModel::MiniCPM4_0_5B => "minicpm4-0.5b",
|
||||||
|
WhichModel::Qwen2_5vl3B => "qwen2.5vl-3b",
|
||||||
|
WhichModel::Qwen2_5vl7B => "qwen2.5vl-7b",
|
||||||
|
WhichModel::Qwen3_0_6B => "qwen3-0.6b",
|
||||||
|
WhichModel::Qwen3ASR0_6B => "qwen3asr-0.6b",
|
||||||
|
WhichModel::Qwen3ASR1_7B => "qwen3asr-1.7b",
|
||||||
|
WhichModel::Qwen3vl2B => "qwen3vl-2b",
|
||||||
|
WhichModel::Qwen3vl4B => "qwen3vl-4b",
|
||||||
|
WhichModel::Qwen3vl8B => "qwen3vl-8b",
|
||||||
|
WhichModel::Qwen3vl32B => "qwen3vl-32b",
|
||||||
|
WhichModel::DeepSeekOCR => "deepseek-ocr",
|
||||||
|
WhichModel::HunyuanOCR => "hunyuan-ocr",
|
||||||
|
WhichModel::PaddleOCRVL => "paddleocr-vl",
|
||||||
|
WhichModel::RMBG2_0 => "rmbg2.0",
|
||||||
|
WhichModel::VoxCPM => "voxcpm",
|
||||||
|
WhichModel::VoxCPM1_5 => "voxcpm1.5",
|
||||||
|
WhichModel::GlmASRNano2512 => "glm-asr-nano-2512",
|
||||||
|
WhichModel::FunASRNano2512 => "fun-asr-nano-2512",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get the owner/organization name for a model
|
||||||
|
fn which_model_to_owner(which_model: WhichModel) -> &'static str {
|
||||||
|
match which_model {
|
||||||
|
WhichModel::MiniCPM4_0_5B => "OpenBMB",
|
||||||
|
WhichModel::Qwen2_5vl3B | WhichModel::Qwen2_5vl7B => "Qwen",
|
||||||
|
WhichModel::Qwen3_0_6B | WhichModel::Qwen3ASR0_6B | WhichModel::Qwen3ASR1_7B => "Qwen",
|
||||||
|
WhichModel::Qwen3vl2B
|
||||||
|
| WhichModel::Qwen3vl4B
|
||||||
|
| WhichModel::Qwen3vl8B
|
||||||
|
| WhichModel::Qwen3vl32B => "Qwen",
|
||||||
|
WhichModel::DeepSeekOCR => "deepseek-ai",
|
||||||
|
WhichModel::HunyuanOCR => "Tencent-Hunyuan",
|
||||||
|
WhichModel::PaddleOCRVL => "PaddlePaddle",
|
||||||
|
WhichModel::RMBG2_0 => "AI-ModelScope",
|
||||||
|
WhichModel::VoxCPM | WhichModel::VoxCPM1_5 => "OpenBMB",
|
||||||
|
WhichModel::GlmASRNano2512 => "ZhipuAI",
|
||||||
|
WhichModel::FunASRNano2512 => "FunAudioLLM",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[get("/models")]
|
||||||
|
pub(crate) async fn models() -> (Status, (ContentType, Json<serde_json::Value>)) {
|
||||||
|
if let Some(model_ref) = MODEL.get() {
|
||||||
|
let guard = model_ref.read().await;
|
||||||
|
let which_model = guard.which_model;
|
||||||
|
|
||||||
|
let model_obj = ModelObject {
|
||||||
|
id: which_model_to_id(which_model).to_string(),
|
||||||
|
object: "model".to_string(),
|
||||||
|
created: None, // We don't track creation time
|
||||||
|
owned_by: which_model_to_owner(which_model).to_string(),
|
||||||
|
};
|
||||||
|
drop(guard);
|
||||||
|
|
||||||
|
let response = ModelsListResponse {
|
||||||
|
object: "list".to_string(),
|
||||||
|
data: vec![model_obj],
|
||||||
|
};
|
||||||
|
(
|
||||||
|
Status::Ok,
|
||||||
|
(
|
||||||
|
ContentType::JSON,
|
||||||
|
Json(serde_json::to_value(response).unwrap()),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
let response = ErrorResponse {
|
||||||
|
error: "model not initialized".to_string(),
|
||||||
|
};
|
||||||
|
(
|
||||||
|
Status::ServiceUnavailable,
|
||||||
|
(
|
||||||
|
ContentType::JSON,
|
||||||
|
Json(serde_json::to_value(response).unwrap()),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
// Test health endpoint when model is not initialized
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_health_endpoint_uninitialized() {
|
||||||
|
let (status, (content_type, response)) = health().await;
|
||||||
|
assert_eq!(status, Status::ServiceUnavailable);
|
||||||
|
assert_eq!(content_type, ContentType::JSON);
|
||||||
|
assert_eq!(response.status, "unhealthy");
|
||||||
|
assert_eq!(response.error, Some("model not initialized".to_string()));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test health endpoint when model is initialized
|
||||||
|
// Note: This test requires a model to be initialized, which may not be feasible
|
||||||
|
// in unit tests without access to model files. This is a placeholder for integration tests.
|
||||||
|
//
|
||||||
|
// #[tokio::test]
|
||||||
|
// async fn test_health_endpoint_initialized() {
|
||||||
|
// // This would require model initialization
|
||||||
|
// // Consider moving to integration tests
|
||||||
|
// }
|
||||||
|
|
||||||
|
// Test models endpoint when model is not initialized
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_models_endpoint_uninitialized() {
|
||||||
|
let (status, (content_type, response)) = models().await;
|
||||||
|
assert_eq!(status, Status::ServiceUnavailable);
|
||||||
|
assert_eq!(content_type, ContentType::JSON);
|
||||||
|
let error = response.get("error").and_then(|v| v.as_str());
|
||||||
|
assert_eq!(error, Some("model not initialized"));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test model type classification
|
||||||
|
#[test]
|
||||||
|
fn test_get_model_type_llm() {
|
||||||
|
assert_eq!(WhichModel::Qwen3_0_6B.model_type(), "llm");
|
||||||
|
assert_eq!(WhichModel::Qwen3vl2B.model_type(), "llm");
|
||||||
|
assert_eq!(WhichModel::MiniCPM4_0_5B.model_type(), "llm");
|
||||||
|
assert_eq!(WhichModel::Qwen2_5vl3B.model_type(), "llm");
|
||||||
|
assert_eq!(WhichModel::Qwen2_5vl7B.model_type(), "llm");
|
||||||
|
assert_eq!(WhichModel::Qwen3vl4B.model_type(), "llm");
|
||||||
|
assert_eq!(WhichModel::Qwen3vl8B.model_type(), "llm");
|
||||||
|
assert_eq!(WhichModel::Qwen3vl32B.model_type(), "llm");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_get_model_type_ocr() {
|
||||||
|
assert_eq!(WhichModel::DeepSeekOCR.model_type(), "ocr");
|
||||||
|
assert_eq!(WhichModel::HunyuanOCR.model_type(), "ocr");
|
||||||
|
assert_eq!(WhichModel::PaddleOCRVL.model_type(), "ocr");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_get_model_type_asr() {
|
||||||
|
assert_eq!(WhichModel::Qwen3ASR0_6B.model_type(), "asr");
|
||||||
|
assert_eq!(WhichModel::Qwen3ASR1_7B.model_type(), "asr");
|
||||||
|
assert_eq!(WhichModel::GlmASRNano2512.model_type(), "asr");
|
||||||
|
assert_eq!(WhichModel::FunASRNano2512.model_type(), "asr");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_get_model_type_image() {
|
||||||
|
assert_eq!(WhichModel::RMBG2_0.model_type(), "image");
|
||||||
|
assert_eq!(WhichModel::VoxCPM.model_type(), "image");
|
||||||
|
assert_eq!(WhichModel::VoxCPM1_5.model_type(), "image");
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test model_id retrieval
|
||||||
|
#[test]
|
||||||
|
fn test_get_model_id() {
|
||||||
|
assert_eq!(WhichModel::Qwen3_0_6B.model_id(), "Qwen/Qwen3-0.6B");
|
||||||
|
assert_eq!(
|
||||||
|
WhichModel::DeepSeekOCR.model_id(),
|
||||||
|
"deepseek-ai/DeepSeek-OCR"
|
||||||
|
);
|
||||||
|
assert_eq!(WhichModel::VoxCPM1_5.model_id(), "OpenBMB/VoxCPM1.5");
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test OpenAI-compatible model ID conversion
|
||||||
|
#[test]
|
||||||
|
fn test_which_model_to_id() {
|
||||||
|
assert_eq!(which_model_to_id(WhichModel::Qwen3_0_6B), "qwen3-0.6b");
|
||||||
|
assert_eq!(which_model_to_id(WhichModel::DeepSeekOCR), "deepseek-ocr");
|
||||||
|
assert_eq!(which_model_to_id(WhichModel::VoxCPM1_5), "voxcpm1.5");
|
||||||
|
assert_eq!(
|
||||||
|
which_model_to_id(WhichModel::MiniCPM4_0_5B),
|
||||||
|
"minicpm4-0.5b"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test owner/organization mapping
|
||||||
|
#[test]
|
||||||
|
fn test_which_model_to_owner() {
|
||||||
|
assert_eq!(which_model_to_owner(WhichModel::Qwen3_0_6B), "Qwen");
|
||||||
|
assert_eq!(which_model_to_owner(WhichModel::DeepSeekOCR), "deepseek-ai");
|
||||||
|
assert_eq!(which_model_to_owner(WhichModel::VoxCPM1_5), "OpenBMB");
|
||||||
|
assert_eq!(
|
||||||
|
which_model_to_owner(WhichModel::HunyuanOCR),
|
||||||
|
"Tencent-Hunyuan"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Shutdown endpoint
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct ShutdownResponse {
|
||||||
|
message: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[post("/shutdown")]
|
||||||
|
pub(crate) async fn shutdown(
|
||||||
|
shutdown_flag: &State<Arc<AtomicBool>>,
|
||||||
|
) -> (Status, (ContentType, Json<serde_json::Value>)) {
|
||||||
|
// Check if remote shutdown is allowed
|
||||||
|
let allow_remote = ALLOW_REMOTE_SHUTDOWN.get().copied().unwrap_or(false);
|
||||||
|
|
||||||
|
// Log the shutdown request
|
||||||
|
eprintln!(
|
||||||
|
"[SHUTDOWN] Shutdown requested (remote_allowed: {})",
|
||||||
|
allow_remote
|
||||||
|
);
|
||||||
|
|
||||||
|
// Note: Rocket 0.5 doesn't provide easy access to client IP in request guards
|
||||||
|
// For proper IP-based filtering, you would need to use custom request guards
|
||||||
|
// or middleware. For now, we rely on the --allow-remote-shutdown flag.
|
||||||
|
|
||||||
|
shutdown_flag.store(true, Ordering::SeqCst);
|
||||||
|
|
||||||
|
// Cleanup PID file in a background task
|
||||||
|
if let Some(&port) = SERVER_PORT.get() {
|
||||||
|
let _ = cleanup_pid_file(port);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Schedule shutdown after a short delay to allow response to be sent
|
||||||
|
let _flag = shutdown_flag.inner().clone();
|
||||||
|
tokio::spawn(async move {
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
||||||
|
std::process::exit(0);
|
||||||
|
});
|
||||||
|
|
||||||
|
let response = ShutdownResponse {
|
||||||
|
message: "Shutting down...".to_string(),
|
||||||
|
};
|
||||||
|
(
|
||||||
|
Status::Ok,
|
||||||
|
(
|
||||||
|
ContentType::JSON,
|
||||||
|
Json(serde_json::to_value(response).unwrap()),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,5 +2,6 @@ pub mod chat_template;
|
|||||||
pub mod exec;
|
pub mod exec;
|
||||||
pub mod models;
|
pub mod models;
|
||||||
pub mod position_embed;
|
pub mod position_embed;
|
||||||
|
pub mod process;
|
||||||
pub mod tokenizer;
|
pub mod tokenizer;
|
||||||
pub mod utils;
|
pub mod utils;
|
||||||
|
|||||||
+251
-45
@@ -1,7 +1,9 @@
|
|||||||
use std::{net::IpAddr, str::FromStr};
|
use std::sync::atomic::{AtomicBool, Ordering};
|
||||||
|
use std::{net::IpAddr, str::FromStr, sync::Arc};
|
||||||
|
|
||||||
use aha::{
|
use aha::{
|
||||||
models::WhichModel,
|
models::WhichModel,
|
||||||
|
process::{cleanup_pid_file, create_pid_file},
|
||||||
utils::{download_model, get_default_save_dir},
|
utils::{download_model, get_default_save_dir},
|
||||||
};
|
};
|
||||||
use clap::{Args, Parser, Subcommand, ValueEnum};
|
use clap::{Args, Parser, Subcommand, ValueEnum};
|
||||||
@@ -10,8 +12,9 @@ use rocket::{
|
|||||||
data::{ByteUnit, Limits},
|
data::{ByteUnit, Limits},
|
||||||
routes,
|
routes,
|
||||||
};
|
};
|
||||||
|
use serde::Serialize;
|
||||||
|
|
||||||
use crate::api::init;
|
use crate::api::{init, set_server_port};
|
||||||
mod api;
|
mod api;
|
||||||
|
|
||||||
#[derive(Parser, Debug)]
|
#[derive(Parser, Debug)]
|
||||||
@@ -52,12 +55,16 @@ enum Commands {
|
|||||||
Cli(CliArgs),
|
Cli(CliArgs),
|
||||||
/// Start service only (--weight-path is optional, defaults to ~/.aha/{model_id})
|
/// Start service only (--weight-path is optional, defaults to ~/.aha/{model_id})
|
||||||
Serv(ServArgs),
|
Serv(ServArgs),
|
||||||
|
/// List all running aha services
|
||||||
|
Ps(ServListArgs),
|
||||||
|
/// Delete a downloaded model from the default location (~/.aha/{model_id})
|
||||||
|
Delete(DeleteArgs),
|
||||||
/// Download model only
|
/// Download model only
|
||||||
Download(DownloadArgs),
|
Download(DownloadArgs),
|
||||||
/// Run model inference directly
|
/// Run model inference directly
|
||||||
Run(RunArgs),
|
Run(RunArgs),
|
||||||
/// List all supported models
|
/// List all supported models
|
||||||
List,
|
List(ListArgs),
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Common/shared arguments for server operations
|
/// Common/shared arguments for server operations
|
||||||
@@ -74,6 +81,10 @@ struct CommonArgs {
|
|||||||
/// Model type (required)
|
/// Model type (required)
|
||||||
#[arg(short, long)]
|
#[arg(short, long)]
|
||||||
model: WhichModel,
|
model: WhichModel,
|
||||||
|
|
||||||
|
/// Allow remote shutdown requests (default: local only, use with caution)
|
||||||
|
#[arg(long)]
|
||||||
|
allow_remote_shutdown: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Arguments for the 'cli' subcommand (download + serve)
|
/// Arguments for the 'cli' subcommand (download + serve)
|
||||||
@@ -95,7 +106,7 @@ struct CliArgs {
|
|||||||
download_retries: Option<u32>,
|
download_retries: Option<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Arguments for the 'serv' subcommand (serve only)
|
/// Arguments for the 'serv start' subcommand
|
||||||
#[derive(Args, Debug)]
|
#[derive(Args, Debug)]
|
||||||
struct ServArgs {
|
struct ServArgs {
|
||||||
#[command(flatten)]
|
#[command(flatten)]
|
||||||
@@ -106,6 +117,14 @@ struct ServArgs {
|
|||||||
weight_path: Option<String>,
|
weight_path: Option<String>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Arguments for the 'serv list' subcommand
|
||||||
|
#[derive(Args, Debug)]
|
||||||
|
struct ServListArgs {
|
||||||
|
/// Compact output format
|
||||||
|
#[arg(short, long)]
|
||||||
|
compact: bool,
|
||||||
|
}
|
||||||
|
|
||||||
/// Arguments for the 'download' subcommand (download only)
|
/// Arguments for the 'download' subcommand (download only)
|
||||||
#[derive(Args, Debug)]
|
#[derive(Args, Debug)]
|
||||||
struct DownloadArgs {
|
struct DownloadArgs {
|
||||||
@@ -142,40 +161,41 @@ struct RunArgs {
|
|||||||
weight_path: Option<String>,
|
weight_path: Option<String>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Arguments for the 'delete' subcommand (delete model from default location)
|
||||||
|
#[derive(Args, Debug)]
|
||||||
|
struct DeleteArgs {
|
||||||
|
/// Model type (required)
|
||||||
|
#[arg(short, long)]
|
||||||
|
model: WhichModel,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Arguments for the 'list' subcommand (list all supported models)
|
||||||
|
#[derive(Args, Debug)]
|
||||||
|
struct ListArgs {
|
||||||
|
/// Output models in JSON format (includes name, model_id, and type fields)
|
||||||
|
#[arg(short, long)]
|
||||||
|
json: bool,
|
||||||
|
}
|
||||||
|
|
||||||
/// Get the default weight path for a given model
|
/// Get the default weight path for a given model
|
||||||
/// Returns ~/.aha/{model_id} e.g., ~/.aha/OpenBMB/VoxCPM1.5
|
/// Returns ~/.aha/{model_id} e.g., ~/.aha/OpenBMB/VoxCPM1.5
|
||||||
fn get_default_weight_path(model: WhichModel) -> String {
|
fn get_default_weight_path(model: WhichModel) -> String {
|
||||||
let model_id = get_model_id(model);
|
let model_id = model.model_id();
|
||||||
let save_dir = get_default_save_dir().expect("Failed to get home directory");
|
let save_dir = get_default_save_dir().expect("Failed to get home directory");
|
||||||
format!("{}/{}", save_dir, model_id)
|
format!("{}/{}", save_dir, model_id)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Get the ModelScope model ID for a given WhichModel variant
|
/// Model information for JSON output
|
||||||
fn get_model_id(model: WhichModel) -> &'static str {
|
#[derive(Serialize)]
|
||||||
match model {
|
struct ModelInfo {
|
||||||
WhichModel::MiniCPM4_0_5B => "OpenBMB/MiniCPM4-0.5B",
|
name: String,
|
||||||
WhichModel::Qwen2_5vl3B => "Qwen/Qwen2.5-VL-3B-Instruct",
|
model_id: String,
|
||||||
WhichModel::Qwen2_5vl7B => "Qwen/Qwen2.5-VL-7B-Instruct",
|
#[serde(rename = "type")]
|
||||||
WhichModel::Qwen3_0_6B => "Qwen/Qwen3-0.6B",
|
model_type: String,
|
||||||
WhichModel::Qwen3ASR0_6B => "Qwen/Qwen3-ASR-0.6B",
|
|
||||||
WhichModel::Qwen3ASR1_7B => "Qwen/Qwen3-ASR-1.7B",
|
|
||||||
WhichModel::Qwen3vl2B => "Qwen/Qwen3-VL-2B-Instruct",
|
|
||||||
WhichModel::Qwen3vl4B => "Qwen/Qwen3-VL-4B-Instruct",
|
|
||||||
WhichModel::Qwen3vl8B => "Qwen/Qwen3-VL-8B-Instruct",
|
|
||||||
WhichModel::Qwen3vl32B => "Qwen/Qwen3-VL-32B-Instruct",
|
|
||||||
WhichModel::DeepSeekOCR => "deepseek-ai/DeepSeek-OCR",
|
|
||||||
WhichModel::HunyuanOCR => "Tencent-Hunyuan/HunyuanOCR",
|
|
||||||
WhichModel::PaddleOCRVL => "PaddlePaddle/PaddleOCR-VL",
|
|
||||||
WhichModel::RMBG2_0 => "AI-ModelScope/RMBG-2.0",
|
|
||||||
WhichModel::VoxCPM => "OpenBMB/VoxCPM-0.5B",
|
|
||||||
WhichModel::VoxCPM1_5 => "OpenBMB/VoxCPM1.5",
|
|
||||||
WhichModel::GlmASRNano2512 => "ZhipuAI/GLM-ASR-Nano-2512",
|
|
||||||
WhichModel::FunASRNano2512 => "FunAudioLLM/Fun-ASR-Nano-2512",
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// List all supported models
|
/// List all supported models
|
||||||
fn run_list() -> anyhow::Result<()> {
|
fn run_list(args: ListArgs) -> anyhow::Result<()> {
|
||||||
let models = [
|
let models = [
|
||||||
WhichModel::MiniCPM4_0_5B,
|
WhichModel::MiniCPM4_0_5B,
|
||||||
WhichModel::Qwen2_5vl3B,
|
WhichModel::Qwen2_5vl3B,
|
||||||
@@ -197,15 +217,32 @@ fn run_list() -> anyhow::Result<()> {
|
|||||||
WhichModel::FunASRNano2512,
|
WhichModel::FunASRNano2512,
|
||||||
];
|
];
|
||||||
|
|
||||||
println!("Available models:");
|
if args.json {
|
||||||
println!();
|
// JSON output
|
||||||
println!("{:<30} ModelScope ID", "Model Name");
|
let model_infos: Vec<ModelInfo> = models
|
||||||
println!("{}", "-".repeat(80));
|
.iter()
|
||||||
for model in models {
|
.map(|model| {
|
||||||
let possible_value = model.to_possible_value().unwrap();
|
let possible_value = model.to_possible_value().unwrap();
|
||||||
let name = possible_value.get_name();
|
ModelInfo {
|
||||||
let id = get_model_id(model);
|
name: possible_value.get_name().to_string(),
|
||||||
println!("{:<30} {}", name, id);
|
model_id: model.model_id().to_string(),
|
||||||
|
model_type: model.model_type().to_string(),
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
println!("{}", serde_json::to_string_pretty(&model_infos)?);
|
||||||
|
} else {
|
||||||
|
// Table output (default)
|
||||||
|
println!("Available models:");
|
||||||
|
println!();
|
||||||
|
println!("{:<30} ModelScope ID", "Model Name");
|
||||||
|
println!("{}", "-".repeat(80));
|
||||||
|
for model in models {
|
||||||
|
let possible_value = model.to_possible_value().unwrap();
|
||||||
|
let name = possible_value.get_name();
|
||||||
|
let id = model.model_id();
|
||||||
|
println!("{:<30} {}", name, id);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
@@ -219,7 +256,7 @@ async fn run_cli(args: CliArgs) -> anyhow::Result<()> {
|
|||||||
save_dir,
|
save_dir,
|
||||||
download_retries,
|
download_retries,
|
||||||
} = args;
|
} = args;
|
||||||
let model_id = get_model_id(common.model);
|
let model_id = common.model.model_id();
|
||||||
|
|
||||||
let model_path = match weight_path {
|
let model_path = match weight_path {
|
||||||
Some(path) => path,
|
Some(path) => path,
|
||||||
@@ -235,7 +272,7 @@ async fn run_cli(args: CliArgs) -> anyhow::Result<()> {
|
|||||||
};
|
};
|
||||||
|
|
||||||
init(common.model, model_path)?;
|
init(common.model, model_path)?;
|
||||||
start_http_server(common.address, common.port).await?;
|
start_http_server(common.address, common.port, common.allow_remote_shutdown).await?;
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
@@ -253,7 +290,48 @@ async fn run_serv(args: ServArgs) -> anyhow::Result<()> {
|
|||||||
};
|
};
|
||||||
|
|
||||||
init(common.model, model_path)?;
|
init(common.model, model_path)?;
|
||||||
start_http_server(common.address, common.port).await?;
|
start_http_server(common.address, common.port, common.allow_remote_shutdown).await?;
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Run the 'ps' subcommand: list running AHA services
|
||||||
|
fn run_ps(args: ServListArgs) -> anyhow::Result<()> {
|
||||||
|
use aha::process::find_aha_services;
|
||||||
|
|
||||||
|
let services = find_aha_services()?;
|
||||||
|
|
||||||
|
if services.is_empty() {
|
||||||
|
println!("No aha services found running.");
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
|
||||||
|
if args.compact {
|
||||||
|
// Compact format: one service per line
|
||||||
|
for svc in services {
|
||||||
|
println!("{}", svc.service_id);
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Table format
|
||||||
|
println!(
|
||||||
|
"{:<20} {:<10} {:<20} {:<10} {:<15} {:<10}",
|
||||||
|
"Service ID", "PID", "Model", "Port", "Address", "Status"
|
||||||
|
);
|
||||||
|
println!("{}", "-".repeat(85));
|
||||||
|
|
||||||
|
for svc in services {
|
||||||
|
let model = svc.model.as_deref().unwrap_or("N/A");
|
||||||
|
let status = match svc.status {
|
||||||
|
aha::process::ServiceStatus::Running => "Running",
|
||||||
|
aha::process::ServiceStatus::Stopping => "Stopping",
|
||||||
|
aha::process::ServiceStatus::Unknown => "Unknown",
|
||||||
|
};
|
||||||
|
println!(
|
||||||
|
"{:<20} {:<10} {:<20} {:<10} {:<15} {:<10}",
|
||||||
|
svc.service_id, svc.pid, model, svc.port, svc.address, status,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
@@ -265,7 +343,7 @@ async fn run_download(args: DownloadArgs) -> anyhow::Result<()> {
|
|||||||
save_dir,
|
save_dir,
|
||||||
download_retries,
|
download_retries,
|
||||||
} = args;
|
} = args;
|
||||||
let model_id = get_model_id(model);
|
let model_id = model.model_id();
|
||||||
|
|
||||||
let save_dir = match save_dir {
|
let save_dir = match save_dir {
|
||||||
Some(dir) => dir,
|
Some(dir) => dir,
|
||||||
@@ -373,6 +451,93 @@ fn run_run(args: RunArgs) -> anyhow::Result<()> {
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Run the 'delete' subcommand: delete model from default location
|
||||||
|
fn run_delete(args: DeleteArgs) -> anyhow::Result<()> {
|
||||||
|
let DeleteArgs { model } = args;
|
||||||
|
let model_id = model.model_id();
|
||||||
|
let save_dir = get_default_save_dir().expect("Failed to get home directory");
|
||||||
|
let model_path = format!("{}/{}", save_dir, model_id);
|
||||||
|
|
||||||
|
let path = std::path::Path::new(&model_path);
|
||||||
|
|
||||||
|
if !path.exists() {
|
||||||
|
println!("Model not found: {} does not exist", model_path);
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
|
||||||
|
// Show model info
|
||||||
|
println!("Model ID: {}", model_id);
|
||||||
|
println!("Location: {}", model_path);
|
||||||
|
|
||||||
|
// Calculate size if possible
|
||||||
|
if let Ok(metadata) = std::fs::metadata(path)
|
||||||
|
&& metadata.is_dir()
|
||||||
|
&& let Ok(total_size) = dir_size(path)
|
||||||
|
{
|
||||||
|
println!("Size: {}", bytes_to_human(total_size));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Confirm deletion
|
||||||
|
print!("Are you sure you want to delete this model? (y/N): ");
|
||||||
|
use std::io::Write;
|
||||||
|
std::io::stdout().flush()?;
|
||||||
|
|
||||||
|
let mut input = String::new();
|
||||||
|
std::io::stdin().read_line(&mut input)?;
|
||||||
|
|
||||||
|
let input = input.trim().to_lowercase();
|
||||||
|
if input != "y" && input != "yes" {
|
||||||
|
println!("Deletion cancelled.");
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete the directory
|
||||||
|
std::fs::remove_dir_all(path)?;
|
||||||
|
|
||||||
|
println!("Model deleted successfully: {}", model_path);
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Calculate total size of a directory recursively
|
||||||
|
fn dir_size(path: &std::path::Path) -> anyhow::Result<u64> {
|
||||||
|
let mut total = 0;
|
||||||
|
if path.is_dir() {
|
||||||
|
for entry in std::fs::read_dir(path)? {
|
||||||
|
let entry = entry?;
|
||||||
|
let entry_path = entry.path();
|
||||||
|
if entry_path.is_dir() {
|
||||||
|
total += dir_size(&entry_path)?;
|
||||||
|
} else {
|
||||||
|
total += entry.metadata()?.len();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
total = std::fs::metadata(path)?.len();
|
||||||
|
}
|
||||||
|
Ok(total)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Convert bytes to human readable format
|
||||||
|
fn bytes_to_human(bytes: u64) -> String {
|
||||||
|
const KB: u64 = 1024;
|
||||||
|
const MB: u64 = KB * 1024;
|
||||||
|
const GB: u64 = MB * 1024;
|
||||||
|
const TB: u64 = GB * 1024;
|
||||||
|
|
||||||
|
if bytes >= TB {
|
||||||
|
format!("{:.2} TB", bytes as f64 / TB as f64)
|
||||||
|
} else if bytes >= GB {
|
||||||
|
format!("{:.2} GB", bytes as f64 / GB as f64)
|
||||||
|
} else if bytes >= MB {
|
||||||
|
format!("{:.2} MB", bytes as f64 / MB as f64)
|
||||||
|
} else if bytes >= KB {
|
||||||
|
format!("{:.2} KB", bytes as f64 / KB as f64)
|
||||||
|
} else {
|
||||||
|
format!("{} B", bytes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::main]
|
#[tokio::main]
|
||||||
async fn main() -> anyhow::Result<()> {
|
async fn main() -> anyhow::Result<()> {
|
||||||
let cli = Cli::parse();
|
let cli = Cli::parse();
|
||||||
@@ -380,9 +545,11 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
match cli.command {
|
match cli.command {
|
||||||
Some(Commands::Cli(args)) => run_cli(args).await,
|
Some(Commands::Cli(args)) => run_cli(args).await,
|
||||||
Some(Commands::Serv(args)) => run_serv(args).await,
|
Some(Commands::Serv(args)) => run_serv(args).await,
|
||||||
|
Some(Commands::Ps(args)) => run_ps(args),
|
||||||
|
Some(Commands::Delete(args)) => run_delete(args),
|
||||||
Some(Commands::Download(args)) => run_download(args).await,
|
Some(Commands::Download(args)) => run_download(args).await,
|
||||||
Some(Commands::Run(args)) => run_run(args),
|
Some(Commands::Run(args)) => run_run(args),
|
||||||
Some(Commands::List) => run_list(),
|
Some(Commands::List(args)) => run_list(args),
|
||||||
None => {
|
None => {
|
||||||
// Backward compatibility: when no subcommand is provided, use 'cli' behavior
|
// Backward compatibility: when no subcommand is provided, use 'cli' behavior
|
||||||
let model = cli.model.expect("Model is required (use -m or --model)");
|
let model = cli.model.expect("Model is required (use -m or --model)");
|
||||||
@@ -391,6 +558,7 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
address: cli.address.unwrap_or_else(|| "127.0.0.1".to_string()),
|
address: cli.address.unwrap_or_else(|| "127.0.0.1".to_string()),
|
||||||
port: cli.port.unwrap_or(10100),
|
port: cli.port.unwrap_or(10100),
|
||||||
model,
|
model,
|
||||||
|
allow_remote_shutdown: false,
|
||||||
},
|
},
|
||||||
weight_path: cli.weight_path,
|
weight_path: cli.weight_path,
|
||||||
save_dir: cli.save_dir,
|
save_dir: cli.save_dir,
|
||||||
@@ -401,7 +569,35 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) async fn start_http_server(address: String, port: u16) -> anyhow::Result<()> {
|
pub(crate) async fn start_http_server(
|
||||||
|
address: String,
|
||||||
|
port: u16,
|
||||||
|
allow_remote_shutdown: bool,
|
||||||
|
) -> anyhow::Result<()> {
|
||||||
|
// Set server port for shutdown endpoint
|
||||||
|
set_server_port(port, allow_remote_shutdown);
|
||||||
|
|
||||||
|
// Create PID file for service tracking
|
||||||
|
let pid = std::process::id();
|
||||||
|
create_pid_file(pid, port)?;
|
||||||
|
|
||||||
|
// Set up shutdown flag
|
||||||
|
let shutdown_flag = Arc::new(AtomicBool::new(false));
|
||||||
|
let shutdown_flag_clone = shutdown_flag.clone();
|
||||||
|
|
||||||
|
// Configure Ctrl+C handler for graceful shutdown
|
||||||
|
let port_for_cleanup = port;
|
||||||
|
let shutdown_handler = tokio::spawn(async move {
|
||||||
|
tokio::signal::ctrl_c().await.ok();
|
||||||
|
println!("Received shutdown signal, gracefully shutting down...");
|
||||||
|
shutdown_flag_clone.store(true, Ordering::SeqCst);
|
||||||
|
// Give time for existing requests to complete
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
||||||
|
// Cleanup PID file
|
||||||
|
let _ = cleanup_pid_file(port_for_cleanup);
|
||||||
|
std::process::exit(0);
|
||||||
|
});
|
||||||
|
|
||||||
let mut builder = rocket::build().configure(Config {
|
let mut builder = rocket::build().configure(Config {
|
||||||
address: IpAddr::from_str(&address)?,
|
address: IpAddr::from_str(&address)?,
|
||||||
port,
|
port,
|
||||||
@@ -416,9 +612,19 @@ pub(crate) async fn start_http_server(address: String, port: u16) -> anyhow::Res
|
|||||||
builder = builder.mount("/chat", routes![api::chat]);
|
builder = builder.mount("/chat", routes![api::chat]);
|
||||||
// /images/remove_background
|
// /images/remove_background
|
||||||
builder = builder.mount("/images", routes![api::remove_background]);
|
builder = builder.mount("/images", routes![api::remove_background]);
|
||||||
// /images/speech
|
// /audio/speech
|
||||||
builder = builder.mount("/audio", routes![api::speech]);
|
builder = builder.mount("/audio", routes![api::speech]);
|
||||||
|
// Health check and model info endpoints
|
||||||
|
builder = builder.mount("/", routes![api::health, api::models]);
|
||||||
|
// Shutdown endpoint
|
||||||
|
builder = builder.manage(shutdown_flag);
|
||||||
|
builder = builder.mount("/", routes![api::shutdown]);
|
||||||
|
|
||||||
|
let _rocket = builder.launch().await?;
|
||||||
|
|
||||||
|
// Cleanup PID file when server exits
|
||||||
|
cleanup_pid_file(port)?;
|
||||||
|
shutdown_handler.abort();
|
||||||
|
|
||||||
builder.launch().await?;
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -74,6 +74,56 @@ pub enum WhichModel {
|
|||||||
FunASRNano2512,
|
FunASRNano2512,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl WhichModel {
|
||||||
|
/// Get the ModelScope model ID for this model variant
|
||||||
|
pub fn model_id(self) -> &'static str {
|
||||||
|
match self {
|
||||||
|
WhichModel::MiniCPM4_0_5B => "OpenBMB/MiniCPM4-0.5B",
|
||||||
|
WhichModel::Qwen2_5vl3B => "Qwen/Qwen2.5-VL-3B-Instruct",
|
||||||
|
WhichModel::Qwen2_5vl7B => "Qwen/Qwen2.5-VL-7B-Instruct",
|
||||||
|
WhichModel::Qwen3_0_6B => "Qwen/Qwen3-0.6B",
|
||||||
|
WhichModel::Qwen3ASR0_6B => "Qwen/Qwen3-ASR-0.6B",
|
||||||
|
WhichModel::Qwen3ASR1_7B => "Qwen/Qwen3-ASR-1.7B",
|
||||||
|
WhichModel::Qwen3vl2B => "Qwen/Qwen3-VL-2B-Instruct",
|
||||||
|
WhichModel::Qwen3vl4B => "Qwen/Qwen3-VL-4B-Instruct",
|
||||||
|
WhichModel::Qwen3vl8B => "Qwen/Qwen3-VL-8B-Instruct",
|
||||||
|
WhichModel::Qwen3vl32B => "Qwen/Qwen3-VL-32B-Instruct",
|
||||||
|
WhichModel::DeepSeekOCR => "deepseek-ai/DeepSeek-OCR",
|
||||||
|
WhichModel::HunyuanOCR => "Tencent-Hunyuan/HunyuanOCR",
|
||||||
|
WhichModel::PaddleOCRVL => "PaddlePaddle/PaddleOCR-VL",
|
||||||
|
WhichModel::RMBG2_0 => "AI-ModelScope/RMBG-2.0",
|
||||||
|
WhichModel::VoxCPM => "OpenBMB/VoxCPM-0.5B",
|
||||||
|
WhichModel::VoxCPM1_5 => "OpenBMB/VoxCPM1.5",
|
||||||
|
WhichModel::GlmASRNano2512 => "ZhipuAI/GLM-ASR-Nano-2512",
|
||||||
|
WhichModel::FunASRNano2512 => "FunAudioLLM/Fun-ASR-Nano-2512",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get the model type category for this model variant
|
||||||
|
pub fn model_type(self) -> &'static str {
|
||||||
|
match self {
|
||||||
|
// LLM models
|
||||||
|
WhichModel::MiniCPM4_0_5B
|
||||||
|
| WhichModel::Qwen2_5vl3B
|
||||||
|
| WhichModel::Qwen2_5vl7B
|
||||||
|
| WhichModel::Qwen3_0_6B
|
||||||
|
| WhichModel::Qwen3vl2B
|
||||||
|
| WhichModel::Qwen3vl4B
|
||||||
|
| WhichModel::Qwen3vl8B
|
||||||
|
| WhichModel::Qwen3vl32B => "llm",
|
||||||
|
// OCR models
|
||||||
|
WhichModel::DeepSeekOCR | WhichModel::HunyuanOCR | WhichModel::PaddleOCRVL => "ocr",
|
||||||
|
// ASR models
|
||||||
|
WhichModel::Qwen3ASR0_6B
|
||||||
|
| WhichModel::Qwen3ASR1_7B
|
||||||
|
| WhichModel::GlmASRNano2512
|
||||||
|
| WhichModel::FunASRNano2512 => "asr",
|
||||||
|
// Image models
|
||||||
|
WhichModel::RMBG2_0 | WhichModel::VoxCPM | WhichModel::VoxCPM1_5 => "image",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
pub trait GenerateModel {
|
pub trait GenerateModel {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse>;
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse>;
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
|
|||||||
+288
@@ -0,0 +1,288 @@
|
|||||||
|
//! Process management module for AHA services
|
||||||
|
//!
|
||||||
|
//! This module provides functionality for:
|
||||||
|
//! - Managing PID files for service tracking
|
||||||
|
//! - Discovering running AHA services
|
||||||
|
//! - Service information display
|
||||||
|
|
||||||
|
use std::fs;
|
||||||
|
use std::path::PathBuf;
|
||||||
|
|
||||||
|
use anyhow::{Result, anyhow};
|
||||||
|
use sysinfo::{Pid, ProcessesToUpdate, System};
|
||||||
|
|
||||||
|
/// Service information structure
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub struct ServiceInfo {
|
||||||
|
/// Service unique identifier (format: pid@port)
|
||||||
|
pub service_id: String,
|
||||||
|
/// Process ID
|
||||||
|
pub pid: u32,
|
||||||
|
/// Model name (if available)
|
||||||
|
pub model: Option<String>,
|
||||||
|
/// Listen port
|
||||||
|
pub port: u16,
|
||||||
|
/// Listen address
|
||||||
|
pub address: String,
|
||||||
|
/// Service status
|
||||||
|
pub status: ServiceStatus,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Service status
|
||||||
|
#[derive(Debug, Clone, PartialEq)]
|
||||||
|
pub enum ServiceStatus {
|
||||||
|
Running,
|
||||||
|
Stopping,
|
||||||
|
Unknown,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get the PID file directory
|
||||||
|
///
|
||||||
|
/// Returns the appropriate directory for storing PID files:
|
||||||
|
/// - Linux/macOS: $XDG_RUNTIME_DIR/aha or ~/.aha/run
|
||||||
|
/// - Windows: %LOCALAPPDATA%\aha\run
|
||||||
|
pub fn get_pid_dir() -> Result<PathBuf> {
|
||||||
|
#[cfg(unix)]
|
||||||
|
{
|
||||||
|
// Try XDG_RUNTIME_DIR first
|
||||||
|
if let Ok(runtime_dir) = std::env::var("XDG_RUNTIME_DIR") {
|
||||||
|
let pid_dir = PathBuf::from(runtime_dir).join("aha");
|
||||||
|
fs::create_dir_all(&pid_dir)?;
|
||||||
|
return Ok(pid_dir);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback to ~/.aha/run
|
||||||
|
let home = dirs::home_dir().ok_or_else(|| anyhow!("Cannot determine home directory"))?;
|
||||||
|
let pid_dir = home.join(".aha").join("run");
|
||||||
|
fs::create_dir_all(&pid_dir)?;
|
||||||
|
Ok(pid_dir)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(windows)]
|
||||||
|
{
|
||||||
|
let local_app_data = std::env::var("LOCALAPPDATA")
|
||||||
|
.map_err(|_| anyhow!("Cannot determine LOCALAPPDATA directory"))?;
|
||||||
|
let pid_dir = PathBuf::from(local_app_data).join("aha").join("run");
|
||||||
|
fs::create_dir_all(&pid_dir)?;
|
||||||
|
Ok(pid_dir)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Create a PID file for the current service
|
||||||
|
///
|
||||||
|
/// # Arguments
|
||||||
|
/// * `pid` - Process ID
|
||||||
|
/// * `port` - Listen port
|
||||||
|
pub fn create_pid_file(pid: u32, port: u16) -> Result<()> {
|
||||||
|
let pid_dir = get_pid_dir()?;
|
||||||
|
let pid_file = pid_dir.join(format!("{}.pid", port));
|
||||||
|
|
||||||
|
let content = format!("{}\n", pid);
|
||||||
|
fs::write(&pid_file, content)?;
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Clean up a PID file
|
||||||
|
///
|
||||||
|
/// # Arguments
|
||||||
|
/// * `port` - Listen port
|
||||||
|
pub fn cleanup_pid_file(port: u16) -> Result<()> {
|
||||||
|
let pid_dir = get_pid_dir()?;
|
||||||
|
let pid_file = pid_dir.join(format!("{}.pid", port));
|
||||||
|
|
||||||
|
if pid_file.exists() {
|
||||||
|
fs::remove_file(&pid_file)?;
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get the PID from a PID file
|
||||||
|
///
|
||||||
|
/// # Arguments
|
||||||
|
/// * `port` - Listen port
|
||||||
|
pub fn get_pid_from_file(port: u16) -> Option<u32> {
|
||||||
|
let pid_dir = get_pid_dir().ok()?;
|
||||||
|
let pid_file = pid_dir.join(format!("{}.pid", port));
|
||||||
|
|
||||||
|
if !pid_file.exists() {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
|
||||||
|
let content = fs::read_to_string(&pid_file).ok()?;
|
||||||
|
content.trim().parse::<u32>().ok()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if a process is an AHA service
|
||||||
|
///
|
||||||
|
/// Verifies that the process command line contains "aha serv" or "aha cli"
|
||||||
|
fn is_aha_process(sys: &System, pid: Pid) -> bool {
|
||||||
|
if let Some(process) = sys.process(pid) {
|
||||||
|
let cmd = process.cmd();
|
||||||
|
let cmd_str: String = cmd
|
||||||
|
.iter()
|
||||||
|
.filter_map(|s| s.to_str())
|
||||||
|
.collect::<Vec<&str>>()
|
||||||
|
.join(" ");
|
||||||
|
return cmd_str.contains("aha serv") || cmd_str.contains("aha cli");
|
||||||
|
}
|
||||||
|
false
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Find all running AHA services
|
||||||
|
///
|
||||||
|
/// Returns a list of ServiceInfo for all running AHA services
|
||||||
|
pub fn find_aha_services() -> Result<Vec<ServiceInfo>> {
|
||||||
|
let mut services = Vec::new();
|
||||||
|
let mut sys = System::new_all();
|
||||||
|
sys.refresh_processes(ProcessesToUpdate::All, true);
|
||||||
|
|
||||||
|
// First, try to discover services from PID files
|
||||||
|
let pid_dir = get_pid_dir()?;
|
||||||
|
if let Ok(entries) = fs::read_dir(&pid_dir) {
|
||||||
|
for entry in entries.flatten() {
|
||||||
|
let path = entry.path();
|
||||||
|
if path.extension().and_then(|s| s.to_str()) != Some("pid") {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract port from filename
|
||||||
|
let port_str = path.file_stem().and_then(|s| s.to_str()).unwrap_or("");
|
||||||
|
let port: u16 = port_str.parse().unwrap_or(0);
|
||||||
|
|
||||||
|
if port == 0 {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read PID from file
|
||||||
|
if let Ok(content) = fs::read_to_string(&path)
|
||||||
|
&& let Ok(pid) = content.trim().parse::<u32>()
|
||||||
|
{
|
||||||
|
let sys_pid = Pid::from_u32(pid);
|
||||||
|
if is_aha_process(&sys, sys_pid) {
|
||||||
|
services.push(ServiceInfo {
|
||||||
|
service_id: format!("{}@{}", pid, port),
|
||||||
|
pid,
|
||||||
|
model: None, // TODO: Extract from command line
|
||||||
|
port,
|
||||||
|
address: "127.0.0.1".to_string(),
|
||||||
|
status: ServiceStatus::Running,
|
||||||
|
});
|
||||||
|
} else {
|
||||||
|
// Stale PID file, remove it
|
||||||
|
let _ = fs::remove_file(&path);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback: scan processes for AHA services
|
||||||
|
for (pid, process) in sys.processes() {
|
||||||
|
if services.iter().any(|s| s.pid == pid.as_u32()) {
|
||||||
|
continue; // Already found via PID file
|
||||||
|
}
|
||||||
|
|
||||||
|
let cmd = process.cmd();
|
||||||
|
let cmd_str: String = cmd
|
||||||
|
.iter()
|
||||||
|
.filter_map(|s| s.to_str())
|
||||||
|
.collect::<Vec<&str>>()
|
||||||
|
.join(" ");
|
||||||
|
|
||||||
|
if cmd_str.contains("aha serv") || cmd_str.contains("aha cli") {
|
||||||
|
// Try to extract port from command line
|
||||||
|
let port_str = cmd
|
||||||
|
.iter()
|
||||||
|
.position(|s| s.to_str() == Some("--port"))
|
||||||
|
.and_then(|i| cmd.get(i + 1))
|
||||||
|
.and_then(|s| s.to_str());
|
||||||
|
let port = port_str
|
||||||
|
.and_then(|s| s.parse::<u16>().ok())
|
||||||
|
.unwrap_or(10100);
|
||||||
|
|
||||||
|
services.push(ServiceInfo {
|
||||||
|
service_id: format!("{}@{}", pid.as_u32(), port),
|
||||||
|
pid: pid.as_u32(),
|
||||||
|
model: None,
|
||||||
|
port,
|
||||||
|
address: "127.0.0.1".to_string(),
|
||||||
|
status: ServiceStatus::Running,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(services)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_get_pid_dir() {
|
||||||
|
let pid_dir = get_pid_dir();
|
||||||
|
assert!(pid_dir.is_ok());
|
||||||
|
let dir = pid_dir.unwrap();
|
||||||
|
assert!(dir.exists());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_create_and_cleanup_pid_file() {
|
||||||
|
let port = 19999;
|
||||||
|
create_pid_file(12345, port).unwrap();
|
||||||
|
let pid = get_pid_from_file(port);
|
||||||
|
assert_eq!(pid, Some(12345));
|
||||||
|
cleanup_pid_file(port).unwrap();
|
||||||
|
let pid = get_pid_from_file(port);
|
||||||
|
assert_eq!(pid, None);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_get_pid_from_file_nonexistent() {
|
||||||
|
let port = 19998; // Use a port that likely doesn't have a PID file
|
||||||
|
let pid = get_pid_from_file(port);
|
||||||
|
assert_eq!(pid, None);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_service_status_debug() {
|
||||||
|
// Test ServiceStatus Debug implementation
|
||||||
|
assert_eq!(format!("{:?}", ServiceStatus::Running), "Running");
|
||||||
|
assert_eq!(format!("{:?}", ServiceStatus::Stopping), "Stopping");
|
||||||
|
assert_eq!(format!("{:?}", ServiceStatus::Unknown), "Unknown");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_service_info_clone() {
|
||||||
|
let service = ServiceInfo {
|
||||||
|
service_id: "12345@10100".to_string(),
|
||||||
|
pid: 12345,
|
||||||
|
model: Some("qwen3-0.6b".to_string()),
|
||||||
|
port: 10100,
|
||||||
|
address: "127.0.0.1".to_string(),
|
||||||
|
status: ServiceStatus::Running,
|
||||||
|
};
|
||||||
|
let service_clone = service.clone();
|
||||||
|
assert_eq!(service_clone.service_id, "12345@10100");
|
||||||
|
assert_eq!(service_clone.pid, 12345);
|
||||||
|
assert_eq!(service_clone.model, Some("qwen3-0.6b".to_string()));
|
||||||
|
assert_eq!(service_clone.port, 10100);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_find_aha_services() {
|
||||||
|
// This test will find actual running AHA services or return empty
|
||||||
|
let services = find_aha_services();
|
||||||
|
assert!(services.is_ok());
|
||||||
|
let services_list = services.unwrap();
|
||||||
|
// We can't assert specific services here since it depends on what's running
|
||||||
|
// but we can verify the structure is correct
|
||||||
|
for service in services_list {
|
||||||
|
assert!(!service.service_id.is_empty());
|
||||||
|
assert!(service.pid > 0);
|
||||||
|
assert!(service.port > 0);
|
||||||
|
assert!(!service.address.is_empty());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+10
-10
@@ -1,17 +1,17 @@
|
|||||||
// use std::io::Cursor;
|
// use std::io::Cursor;
|
||||||
|
|
||||||
use std::fs::File;
|
// use std::fs::File;
|
||||||
// use symphonia::core::io::MediaSourceStream;
|
// use symphonia::core::io::MediaSourceStream;
|
||||||
use std::io::{Read, Seek};
|
// use std::io::{Read, Seek};
|
||||||
use std::{io::Cursor, time::Instant};
|
// use std::{io::Cursor, time::Instant};
|
||||||
|
|
||||||
use aha::utils::{load_tensor_from_pt, tensor_utils::interpolate_nearest_1d};
|
use aha::utils::load_tensor_from_pt;
|
||||||
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
// use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::Result;
|
||||||
use byteorder::{LittleEndian, ReadBytesExt};
|
// use byteorder::{LittleEndian, ReadBytesExt};
|
||||||
use candle_core::{Shape, Tensor};
|
use candle_core::Shape;
|
||||||
use sentencepiece::SentencePieceProcessor;
|
// use sentencepiece::SentencePieceProcessor;
|
||||||
use zip::ZipArchive;
|
// use zip::ZipArchive;
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn messy_test() -> Result<()> {
|
fn messy_test() -> Result<()> {
|
||||||
|
|||||||
@@ -0,0 +1,76 @@
|
|||||||
|
use aha::models::WhichModel;
|
||||||
|
|
||||||
|
// Import helper functions from api module - these will need to be made public
|
||||||
|
// or tested through integration testing
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_model_type_classification() {
|
||||||
|
// Since get_model_type and get_model_id are private to api.rs,
|
||||||
|
// we document the expected behavior here for reference:
|
||||||
|
//
|
||||||
|
// LLM models: MiniCPM4_0_5B, Qwen2_5vl3B, Qwen2_5vl7B, Qwen3_0_6B,
|
||||||
|
// Qwen3vl2B, Qwen3vl4B, Qwen3vl8B, Qwen3vl32B
|
||||||
|
// OCR models: DeepSeekOCR, HunyuanOCR, PaddleOCRVL
|
||||||
|
// ASR models: Qwen3ASR0_6B, Qwen3ASR1_7B, GlmASRNano2512, FunASRNano2512
|
||||||
|
// Image models: RMBG2_0, VoxCPM, VoxCPM1_5
|
||||||
|
|
||||||
|
// This test documents the expected model type classification
|
||||||
|
let llm_models = vec![
|
||||||
|
WhichModel::MiniCPM4_0_5B,
|
||||||
|
WhichModel::Qwen2_5vl3B,
|
||||||
|
WhichModel::Qwen2_5vl7B,
|
||||||
|
WhichModel::Qwen3_0_6B,
|
||||||
|
WhichModel::Qwen3vl2B,
|
||||||
|
WhichModel::Qwen3vl4B,
|
||||||
|
WhichModel::Qwen3vl8B,
|
||||||
|
WhichModel::Qwen3vl32B,
|
||||||
|
];
|
||||||
|
|
||||||
|
let ocr_models = vec![
|
||||||
|
WhichModel::DeepSeekOCR,
|
||||||
|
WhichModel::HunyuanOCR,
|
||||||
|
WhichModel::PaddleOCRVL,
|
||||||
|
];
|
||||||
|
|
||||||
|
let asr_models = vec![
|
||||||
|
WhichModel::Qwen3ASR0_6B,
|
||||||
|
WhichModel::Qwen3ASR1_7B,
|
||||||
|
WhichModel::GlmASRNano2512,
|
||||||
|
WhichModel::FunASRNano2512,
|
||||||
|
];
|
||||||
|
|
||||||
|
let image_models = vec![
|
||||||
|
WhichModel::RMBG2_0,
|
||||||
|
WhichModel::VoxCPM,
|
||||||
|
WhichModel::VoxCPM1_5,
|
||||||
|
];
|
||||||
|
|
||||||
|
// Verify counts
|
||||||
|
assert_eq!(llm_models.len(), 8);
|
||||||
|
assert_eq!(ocr_models.len(), 3);
|
||||||
|
assert_eq!(asr_models.len(), 4);
|
||||||
|
assert_eq!(image_models.len(), 3);
|
||||||
|
|
||||||
|
// Total models
|
||||||
|
assert_eq!(
|
||||||
|
llm_models.len() + ocr_models.len() + asr_models.len() + image_models.len(),
|
||||||
|
18
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Note: Integration tests for the /health and /models endpoints
|
||||||
|
// should be done with a running server. These would typically:
|
||||||
|
//
|
||||||
|
// 1. Start the server with a test model
|
||||||
|
// 2. Make HTTP requests to /health and /models
|
||||||
|
// 3. Verify the response format and status codes
|
||||||
|
//
|
||||||
|
// Example (pseudo-code):
|
||||||
|
//
|
||||||
|
// #[tokio::test]
|
||||||
|
// async fn test_health_endpoint() {
|
||||||
|
// let resp = reqwest::get("http://localhost:10100/health").await.unwrap();
|
||||||
|
// assert_eq!(resp.status(), 200);
|
||||||
|
// let json: serde_json::Value = resp.json().await.unwrap();
|
||||||
|
// assert_eq!(json["status"], "ok");
|
||||||
|
// }
|
||||||
@@ -1,12 +1,8 @@
|
|||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
|
|
||||||
use aha::utils::{find_type_files, get_device, read_pth_tensor_info_cycle};
|
use aha::utils::{find_type_files, get_device};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{
|
use candle_core::{Device, pickle::read_all_with_key, safetensors};
|
||||||
Device,
|
|
||||||
pickle::{read_all_with_key, read_pth_tensor_info},
|
|
||||||
safetensors,
|
|
||||||
};
|
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -207,7 +203,7 @@ fn index_tts2_weight() -> Result<()> {
|
|||||||
// RUST_BACKTRACE=1 cargo test -F cuda index_tts2_weight -r -- --nocapture
|
// RUST_BACKTRACE=1 cargo test -F cuda index_tts2_weight -r -- --nocapture
|
||||||
let save_dir: String =
|
let save_dir: String =
|
||||||
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
|
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
|
||||||
let model_path = format!("{}/IndexTeam/IndexTTS-2/", save_dir);
|
// let model_path = format!("{}/IndexTeam/IndexTTS-2/", save_dir);
|
||||||
let bigvgan_path = format!(
|
let bigvgan_path = format!(
|
||||||
"{}/nv-community/bigvgan_v2_22khz_80band_256x/bigvgan_generator.pt",
|
"{}/nv-community/bigvgan_v2_22khz_80band_256x/bigvgan_generator.pt",
|
||||||
save_dir
|
save_dir
|
||||||
|
|||||||
Reference in New Issue
Block a user