Merge branch 'main' into self_upgrade
This commit is contained in:
commit
5c5e94c6d8
52 changed files with 1232 additions and 434 deletions
Binary file not shown.
|
Before Width: | Height: | Size: 321 KiB After Width: | Height: | Size: 365 KiB |
|
|
@ -48,6 +48,10 @@
|
||||||
"model": "deepseek/deepseek-chat",
|
"model": "deepseek/deepseek-chat",
|
||||||
"api_key": "sk-your-deepseek-key"
|
"api_key": "sk-your-deepseek-key"
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
"model_name": "lmstudio-local",
|
||||||
|
"model": "lmstudio/openai/gpt-oss-20b"
|
||||||
|
},
|
||||||
{
|
{
|
||||||
"model_name": "longcat",
|
"model_name": "longcat",
|
||||||
"model": "longcat/LongCat-Flash-Thinking",
|
"model": "longcat/LongCat-Flash-Thinking",
|
||||||
|
|
|
||||||
|
|
@ -563,6 +563,7 @@ For complete documentation, see [`security_configuration.md`](security_configura
|
||||||
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [Get Key](https://dashscope.console.aliyun.com) |
|
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [Get Key](https://dashscope.console.aliyun.com) |
|
||||||
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [Get Key](https://build.nvidia.com) |
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [Get Key](https://build.nvidia.com) |
|
||||||
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | Local (no key needed) |
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | Local (no key needed) |
|
||||||
|
| **LM Studio** | `lmstudio/` | `http://localhost:1234/v1` | OpenAI | Optional (local default: no key) |
|
||||||
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [Get Key](https://openrouter.ai/keys) |
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [Get Key](https://openrouter.ai/keys) |
|
||||||
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | Your LiteLLM proxy key |
|
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | Your LiteLLM proxy key |
|
||||||
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | Local |
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | Local |
|
||||||
|
|
@ -710,6 +711,21 @@ For direct Anthropic API access or custom endpoints that only support Anthropic'
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>LM Studio (local)</b></summary>
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "lmstudio-local",
|
||||||
|
"model": "lmstudio/openai/gpt-oss-20b"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
`api_base` defaults to `http://localhost:1234/v1`. API key is optional unless your LM Studio server enables authentication.<br/>
|
||||||
|
PicoClaw sends OpenAI-compatible requests to LM Studio, and strips the `lmstudio/` prefix before sending requests, so `lmstudio/openai/gpt-oss-20b` sends `openai/gpt-oss-20b` to the LM Studio server.
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>Custom Proxy / LiteLLM</b></summary>
|
<summary><b>Custom Proxy / LiteLLM</b></summary>
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -56,6 +56,7 @@ This design also enables **multi-agent support** with flexible provider selectio
|
||||||
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [Get Key](https://dashscope.console.aliyun.com) |
|
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [Get Key](https://dashscope.console.aliyun.com) |
|
||||||
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [Get Key](https://build.nvidia.com) |
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [Get Key](https://build.nvidia.com) |
|
||||||
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | Local (no key needed) |
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | Local (no key needed) |
|
||||||
|
| **LM Studio** | `lmstudio/` | `http://localhost:1234/v1` | OpenAI | Optional (local default: no key) |
|
||||||
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [Get Key](https://openrouter.ai/keys) |
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [Get Key](https://openrouter.ai/keys) |
|
||||||
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | Your LiteLLM proxy key |
|
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | Your LiteLLM proxy key |
|
||||||
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | Local |
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | Local |
|
||||||
|
|
@ -226,6 +227,18 @@ For direct Anthropic API access or custom endpoints that only support Anthropic'
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
**LM Studio (local)**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "lmstudio-local",
|
||||||
|
"model": "lmstudio/openai/gpt-oss-20b"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
`api_base` defaults to `http://localhost:1234/v1`. API key is optional unless your LM Studio server enables authentication.<br/>
|
||||||
|
PicoClaw sends OpenAI-compatible requests to LM Studio, and strips the `lmstudio/` prefix before sending requests, so `lmstudio/openai/gpt-oss-20b` sends `openai/gpt-oss-20b` to the LM Studio server.
|
||||||
|
|
||||||
**Custom Proxy/API**
|
**Custom Proxy/API**
|
||||||
|
|
||||||
```json
|
```json
|
||||||
|
|
|
||||||
|
|
@ -365,6 +365,7 @@ Agent 读取 HEARTBEAT.md
|
||||||
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [获取](https://dashscope.console.aliyun.com) |
|
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [获取](https://dashscope.console.aliyun.com) |
|
||||||
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [获取](https://build.nvidia.com) |
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [获取](https://build.nvidia.com) |
|
||||||
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | 本地(无需 Key) |
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | 本地(无需 Key) |
|
||||||
|
| **LM Studio** | `lmstudio/` | `http://localhost:1234/v1` | OpenAI | 可选(本地默认无需密钥) |
|
||||||
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [获取](https://openrouter.ai/keys) |
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [获取](https://openrouter.ai/keys) |
|
||||||
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | 你的 LiteLLM 代理 Key |
|
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | 你的 LiteLLM 代理 Key |
|
||||||
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | 本地 |
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | 本地 |
|
||||||
|
|
@ -506,6 +507,21 @@ Agent 读取 HEARTBEAT.md
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>LM Studio(本地)</b></summary>
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "lmstudio-local",
|
||||||
|
"model": "lmstudio/openai/gpt-oss-20b"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
`api_base` 默认是 `http://localhost:1234/v1`。除非你在 LM Studio 侧启用了认证,否则不需要配置 API Key。
|
||||||
|
PicoClaw 向 LM Studio 的 OpenAI 兼容终结点发送请求,且将移除首个 `lmstudio/` 前缀,因此 `lmstudio/openai/gpt-oss-20b` 会发送 `openai/gpt-oss-20b`。
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>自定义代理 / LiteLLM</b></summary>
|
<summary><b>自定义代理 / LiteLLM</b></summary>
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -53,6 +53,7 @@
|
||||||
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [获取密钥](https://dashscope.console.aliyun.com) |
|
| **通义千问 (Qwen)** | `qwen/` | `https://dashscope.aliyuncs.com/compatible-mode/v1` | OpenAI | [获取密钥](https://dashscope.console.aliyun.com) |
|
||||||
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [获取密钥](https://build.nvidia.com) |
|
| **NVIDIA** | `nvidia/` | `https://integrate.api.nvidia.com/v1` | OpenAI | [获取密钥](https://build.nvidia.com) |
|
||||||
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | 本地(无需密钥) |
|
| **Ollama** | `ollama/` | `http://localhost:11434/v1` | OpenAI | 本地(无需密钥) |
|
||||||
|
| **LM Studio** | `lmstudio/` | `http://localhost:1234/v1` | OpenAI | 可选(本地默认无需密钥) |
|
||||||
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [获取密钥](https://openrouter.ai/keys) |
|
| **OpenRouter** | `openrouter/` | `https://openrouter.ai/api/v1` | OpenAI | [获取密钥](https://openrouter.ai/keys) |
|
||||||
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | 你的 LiteLLM 代理密钥 |
|
| **LiteLLM Proxy** | `litellm/` | `http://localhost:4000/v1` | OpenAI | 你的 LiteLLM 代理密钥 |
|
||||||
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | 本地 |
|
| **VLLM** | `vllm/` | `http://localhost:8000/v1` | OpenAI | 本地 |
|
||||||
|
|
@ -211,6 +212,18 @@
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
**LM Studio(本地)**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_name": "lmstudio-local",
|
||||||
|
"model": "lmstudio/openai/gpt-oss-20b"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
`api_base` 默认是 `http://localhost:1234/v1`。除非你在 LM Studio 侧启用了认证,否则不需要配置 API Key。
|
||||||
|
PicoClaw 向 LM Studio 的 OpenAI 兼容终结点发送请求,且将移除首个 `lmstudio/` 前缀,因此 `lmstudio/openai/gpt-oss-20b` 会发送 `openai/gpt-oss-20b`。
|
||||||
|
|
||||||
**自定义代理/API**
|
**自定义代理/API**
|
||||||
|
|
||||||
```json
|
```json
|
||||||
|
|
|
||||||
|
|
@ -25,23 +25,25 @@ import (
|
||||||
|
|
||||||
type fakeChannel struct{ id string }
|
type fakeChannel struct{ id string }
|
||||||
|
|
||||||
func (f *fakeChannel) Name() string { return "fake" }
|
func (f *fakeChannel) Name() string { return "fake" }
|
||||||
func (f *fakeChannel) Start(ctx context.Context) error { return nil }
|
func (f *fakeChannel) Start(ctx context.Context) error { return nil }
|
||||||
func (f *fakeChannel) Stop(ctx context.Context) error { return nil }
|
func (f *fakeChannel) Stop(ctx context.Context) error { return nil }
|
||||||
func (f *fakeChannel) Send(ctx context.Context, msg bus.OutboundMessage) error { return nil }
|
func (f *fakeChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
func (f *fakeChannel) IsRunning() bool { return true }
|
return nil, nil
|
||||||
func (f *fakeChannel) IsAllowed(string) bool { return true }
|
}
|
||||||
func (f *fakeChannel) IsAllowedSender(sender bus.SenderInfo) bool { return true }
|
func (f *fakeChannel) IsRunning() bool { return true }
|
||||||
func (f *fakeChannel) ReasoningChannelID() string { return f.id }
|
func (f *fakeChannel) IsAllowed(string) bool { return true }
|
||||||
|
func (f *fakeChannel) IsAllowedSender(sender bus.SenderInfo) bool { return true }
|
||||||
|
func (f *fakeChannel) ReasoningChannelID() string { return f.id }
|
||||||
|
|
||||||
type fakeMediaChannel struct {
|
type fakeMediaChannel struct {
|
||||||
fakeChannel
|
fakeChannel
|
||||||
sentMedia []bus.OutboundMediaMessage
|
sentMedia []bus.OutboundMediaMessage
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *fakeMediaChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (f *fakeMediaChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
f.sentMedia = append(f.sentMedia, msg)
|
f.sentMedia = append(f.sentMedia, msg)
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func newStartedTestChannelManager(
|
func newStartedTestChannelManager(
|
||||||
|
|
|
||||||
|
|
@ -252,28 +252,28 @@ func (c *TelegramChannel) Stop(ctx context.Context) error {
|
||||||
**3e. Send method error returns**
|
**3e. Send method error returns**
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// Old code: returns plain error
|
// Old code: returned only error
|
||||||
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
if !c.running { return fmt.Errorf("not running") }
|
if !c.running { return fmt.Errorf("not running") }
|
||||||
// ...
|
// ...
|
||||||
if err != nil { return err }
|
if err != nil { return err }
|
||||||
}
|
}
|
||||||
|
|
||||||
// New code: must return sentinel errors for Manager to determine retry strategy
|
// New code: return delivered message IDs plus sentinel errors
|
||||||
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning // ← Manager will not retry
|
return nil, channels.ErrNotRunning // ← Manager will not retry
|
||||||
}
|
}
|
||||||
// ...
|
// ...
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Use ClassifySendError to wrap error based on HTTP status code
|
// Use ClassifySendError to wrap error based on HTTP status code
|
||||||
return channels.ClassifySendError(statusCode, err)
|
return nil, channels.ClassifySendError(statusCode, err)
|
||||||
// Or manually wrap:
|
// Or manually wrap:
|
||||||
// return fmt.Errorf("%w: %v", channels.ErrTemporary, err)
|
// return nil, fmt.Errorf("%w: %v", channels.ErrTemporary, err)
|
||||||
// return fmt.Errorf("%w: %v", channels.ErrRateLimit, err)
|
// return nil, fmt.Errorf("%w: %v", channels.ErrRateLimit, err)
|
||||||
// return fmt.Errorf("%w: %v", channels.ErrSendFailed, err)
|
// return nil, fmt.Errorf("%w: %v", channels.ErrSendFailed, err)
|
||||||
}
|
}
|
||||||
return nil
|
return []string{deliveredID}, nil // or return nil, nil if IDs are unavailable
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -502,25 +502,25 @@ func (c *MatrixChannel) Stop(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *MatrixChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *MatrixChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
// 1. Check running state
|
// 1. Check running state
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
// 2. Send message to Matrix
|
// 2. Send message to Matrix
|
||||||
err := c.sendToMatrix(ctx, msg.ChatID, msg.Content)
|
eventID, err := c.sendToMatrix(ctx, msg.ChatID, msg.Content)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// 3. Must use error classification wrapping
|
// 3. Must use error classification wrapping
|
||||||
// If you have an HTTP status code:
|
// If you have an HTTP status code:
|
||||||
// return channels.ClassifySendError(statusCode, err)
|
// return nil, channels.ClassifySendError(statusCode, err)
|
||||||
// If it's a network error:
|
// If it's a network error:
|
||||||
// return channels.ClassifyNetError(err)
|
// return nil, channels.ClassifyNetError(err)
|
||||||
// If manual classification is needed:
|
// If manual classification is needed:
|
||||||
return fmt.Errorf("%w: %v", channels.ErrTemporary, err)
|
return nil, fmt.Errorf("%w: %v", channels.ErrTemporary, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return []string{eventID}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ========== Incoming Message Handling ==========
|
// ========== Incoming Message Handling ==========
|
||||||
|
|
@ -580,9 +580,9 @@ func (c *MatrixChannel) handleIncoming(roomID, senderID, displayName, content st
|
||||||
|
|
||||||
// ========== Internal Methods ==========
|
// ========== Internal Methods ==========
|
||||||
|
|
||||||
func (c *MatrixChannel) sendToMatrix(ctx context.Context, roomID, content string) error {
|
func (c *MatrixChannel) sendToMatrix(ctx context.Context, roomID, content string) (string, error) {
|
||||||
// Actual Matrix SDK call
|
// Actual Matrix SDK call
|
||||||
return nil
|
return "event-id", nil
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -594,16 +594,17 @@ Depending on platform capabilities, your channel can optionally implement the fo
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// If the platform supports sending images/files/audio/video
|
// If the platform supports sending images/files/audio/video
|
||||||
func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
store := c.GetMediaStore()
|
store := c.GetMediaStore()
|
||||||
if store == nil {
|
if store == nil {
|
||||||
return fmt.Errorf("no media store: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("no media store: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var messageIDs []string
|
||||||
for _, part := range msg.Parts {
|
for _, part := range msg.Parts {
|
||||||
localPath, err := store.Resolve(part.Ref)
|
localPath, err := store.Resolve(part.Ref)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -620,8 +621,10 @@ func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
default:
|
default:
|
||||||
// Upload file to Matrix
|
// Upload file to Matrix
|
||||||
}
|
}
|
||||||
|
// Append platform IDs here when the API returns them.
|
||||||
|
// messageIDs = append(messageIDs, uploadedMessageID)
|
||||||
}
|
}
|
||||||
return nil
|
return messageIDs, nil
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -1270,7 +1273,7 @@ type Channel interface {
|
||||||
Name() string
|
Name() string
|
||||||
Start(ctx context.Context) error
|
Start(ctx context.Context) error
|
||||||
Stop(ctx context.Context) error
|
Stop(ctx context.Context) error
|
||||||
Send(ctx context.Context, msg bus.OutboundMessage) error
|
Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error)
|
||||||
IsRunning() bool
|
IsRunning() bool
|
||||||
IsAllowed(senderID string) bool
|
IsAllowed(senderID string) bool
|
||||||
IsAllowedSender(sender bus.SenderInfo) bool
|
IsAllowedSender(sender bus.SenderInfo) bool
|
||||||
|
|
@ -1279,7 +1282,7 @@ type Channel interface {
|
||||||
|
|
||||||
// ===== Optional =====
|
// ===== Optional =====
|
||||||
type MediaSender interface {
|
type MediaSender interface {
|
||||||
SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error
|
SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
type TypingCapable interface {
|
type TypingCapable interface {
|
||||||
|
|
|
||||||
|
|
@ -252,28 +252,28 @@ func (c *TelegramChannel) Stop(ctx context.Context) error {
|
||||||
**3e. Send 方法的错误返回**
|
**3e. Send 方法的错误返回**
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// 旧代码:返回普通 error
|
// 旧代码:只返回 error
|
||||||
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
if !c.running { return fmt.Errorf("not running") }
|
if !c.running { return fmt.Errorf("not running") }
|
||||||
// ...
|
// ...
|
||||||
if err != nil { return err }
|
if err != nil { return err }
|
||||||
}
|
}
|
||||||
|
|
||||||
// 新代码:必须返回哨兵错误,供 Manager 判断重试策略
|
// 新代码:返回投递后的消息 ID,以及供 Manager 判断重试策略的哨兵错误
|
||||||
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning // ← Manager 不会重试
|
return nil, channels.ErrNotRunning // ← Manager 不会重试
|
||||||
}
|
}
|
||||||
// ...
|
// ...
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// 使用 ClassifySendError 根据 HTTP 状态码包装错误
|
// 使用 ClassifySendError 根据 HTTP 状态码包装错误
|
||||||
return channels.ClassifySendError(statusCode, err)
|
return nil, channels.ClassifySendError(statusCode, err)
|
||||||
// 或手动包装:
|
// 或手动包装:
|
||||||
// return fmt.Errorf("%w: %v", channels.ErrTemporary, err)
|
// return nil, fmt.Errorf("%w: %v", channels.ErrTemporary, err)
|
||||||
// return fmt.Errorf("%w: %v", channels.ErrRateLimit, err)
|
// return nil, fmt.Errorf("%w: %v", channels.ErrRateLimit, err)
|
||||||
// return fmt.Errorf("%w: %v", channels.ErrSendFailed, err)
|
// return nil, fmt.Errorf("%w: %v", channels.ErrSendFailed, err)
|
||||||
}
|
}
|
||||||
return nil
|
return []string{deliveredID}, nil // 如果拿不到 ID,也可以返回 nil, nil
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -502,25 +502,25 @@ func (c *MatrixChannel) Stop(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *MatrixChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *MatrixChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
// 1. 检查运行状态
|
// 1. 检查运行状态
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
// 2. 发送消息到 Matrix
|
// 2. 发送消息到 Matrix
|
||||||
err := c.sendToMatrix(ctx, msg.ChatID, msg.Content)
|
eventID, err := c.sendToMatrix(ctx, msg.ChatID, msg.Content)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// 3. 必须使用错误分类包装
|
// 3. 必须使用错误分类包装
|
||||||
// 如果你有 HTTP 状态码:
|
// 如果你有 HTTP 状态码:
|
||||||
// return channels.ClassifySendError(statusCode, err)
|
// return nil, channels.ClassifySendError(statusCode, err)
|
||||||
// 如果是网络错误:
|
// 如果是网络错误:
|
||||||
// return channels.ClassifyNetError(err)
|
// return nil, channels.ClassifyNetError(err)
|
||||||
// 如果需要手动分类:
|
// 如果需要手动分类:
|
||||||
return fmt.Errorf("%w: %v", channels.ErrTemporary, err)
|
return nil, fmt.Errorf("%w: %v", channels.ErrTemporary, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return []string{eventID}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ========== 消息接收处理 ==========
|
// ========== 消息接收处理 ==========
|
||||||
|
|
@ -580,9 +580,9 @@ func (c *MatrixChannel) handleIncoming(roomID, senderID, displayName, content st
|
||||||
|
|
||||||
// ========== 内部方法 ==========
|
// ========== 内部方法 ==========
|
||||||
|
|
||||||
func (c *MatrixChannel) sendToMatrix(ctx context.Context, roomID, content string) error {
|
func (c *MatrixChannel) sendToMatrix(ctx context.Context, roomID, content string) (string, error) {
|
||||||
// 实际的 Matrix SDK 调用
|
// 实际的 Matrix SDK 调用
|
||||||
return nil
|
return "event-id", nil
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -594,16 +594,17 @@ func (c *MatrixChannel) sendToMatrix(ctx context.Context, roomID, content string
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// 如果平台支持发送图片/文件/音频/视频
|
// 如果平台支持发送图片/文件/音频/视频
|
||||||
func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
store := c.GetMediaStore()
|
store := c.GetMediaStore()
|
||||||
if store == nil {
|
if store == nil {
|
||||||
return fmt.Errorf("no media store: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("no media store: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var messageIDs []string
|
||||||
for _, part := range msg.Parts {
|
for _, part := range msg.Parts {
|
||||||
localPath, err := store.Resolve(part.Ref)
|
localPath, err := store.Resolve(part.Ref)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -620,8 +621,10 @@ func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
default:
|
default:
|
||||||
// 上传文件到 Matrix
|
// 上传文件到 Matrix
|
||||||
}
|
}
|
||||||
|
// 如果 API 能返回平台消息 ID,就在这里追加。
|
||||||
|
// messageIDs = append(messageIDs, uploadedMessageID)
|
||||||
}
|
}
|
||||||
return nil
|
return messageIDs, nil
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -1269,7 +1272,7 @@ type Channel interface {
|
||||||
Name() string
|
Name() string
|
||||||
Start(ctx context.Context) error
|
Start(ctx context.Context) error
|
||||||
Stop(ctx context.Context) error
|
Stop(ctx context.Context) error
|
||||||
Send(ctx context.Context, msg bus.OutboundMessage) error
|
Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error)
|
||||||
IsRunning() bool
|
IsRunning() bool
|
||||||
IsAllowed(senderID string) bool
|
IsAllowed(senderID string) bool
|
||||||
IsAllowedSender(sender bus.SenderInfo) bool
|
IsAllowedSender(sender bus.SenderInfo) bool
|
||||||
|
|
@ -1278,7 +1281,7 @@ type Channel interface {
|
||||||
|
|
||||||
// ===== 可选实现 =====
|
// ===== 可选实现 =====
|
||||||
type MediaSender interface {
|
type MediaSender interface {
|
||||||
SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error
|
SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
type TypingCapable interface {
|
type TypingCapable interface {
|
||||||
|
|
|
||||||
|
|
@ -48,7 +48,7 @@ type Channel interface {
|
||||||
Name() string
|
Name() string
|
||||||
Start(ctx context.Context) error
|
Start(ctx context.Context) error
|
||||||
Stop(ctx context.Context) error
|
Stop(ctx context.Context) error
|
||||||
Send(ctx context.Context, msg bus.OutboundMessage) error
|
Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error)
|
||||||
IsRunning() bool
|
IsRunning() bool
|
||||||
IsAllowed(senderID string) bool
|
IsAllowed(senderID string) bool
|
||||||
IsAllowedSender(sender bus.SenderInfo) bool
|
IsAllowedSender(sender bus.SenderInfo) bool
|
||||||
|
|
|
||||||
|
|
@ -104,20 +104,20 @@ func (c *DingTalkChannel) Stop(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Send sends a message to DingTalk via the chatbot reply API
|
// Send sends a message to DingTalk via the chatbot reply API
|
||||||
func (c *DingTalkChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *DingTalkChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get session webhook from storage
|
// Get session webhook from storage
|
||||||
sessionWebhookRaw, ok := c.sessionWebhooks.Load(msg.ChatID)
|
sessionWebhookRaw, ok := c.sessionWebhooks.Load(msg.ChatID)
|
||||||
if !ok {
|
if !ok {
|
||||||
return fmt.Errorf("no session_webhook found for chat %s, cannot send message", msg.ChatID)
|
return nil, fmt.Errorf("no session_webhook found for chat %s, cannot send message", msg.ChatID)
|
||||||
}
|
}
|
||||||
|
|
||||||
sessionWebhook, ok := sessionWebhookRaw.(string)
|
sessionWebhook, ok := sessionWebhookRaw.(string)
|
||||||
if !ok {
|
if !ok {
|
||||||
return fmt.Errorf("invalid session_webhook type for chat %s", msg.ChatID)
|
return nil, fmt.Errorf("invalid session_webhook type for chat %s", msg.ChatID)
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.DebugCF("dingtalk", "Sending message", map[string]any{
|
logger.DebugCF("dingtalk", "Sending message", map[string]any{
|
||||||
|
|
@ -126,7 +126,7 @@ func (c *DingTalkChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
})
|
})
|
||||||
|
|
||||||
// Use the session webhook to send the reply
|
// Use the session webhook to send the reply
|
||||||
return c.SendDirectReply(ctx, sessionWebhook, msg.Content)
|
return nil, c.SendDirectReply(ctx, sessionWebhook, msg.Content)
|
||||||
}
|
}
|
||||||
|
|
||||||
// onChatBotMessageReceived implements the IChatBotMessageHandler function signature
|
// onChatBotMessageReceived implements the IChatBotMessageHandler function signature
|
||||||
|
|
|
||||||
|
|
@ -128,37 +128,41 @@ func (c *DiscordChannel) Stop(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
channelID := msg.ChatID
|
channelID := msg.ChatID
|
||||||
if channelID == "" {
|
if channelID == "" {
|
||||||
return fmt.Errorf("channel ID is empty")
|
return nil, fmt.Errorf("channel ID is empty")
|
||||||
}
|
}
|
||||||
|
|
||||||
if len([]rune(msg.Content)) == 0 {
|
if len([]rune(msg.Content)) == 0 {
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return c.sendChunk(ctx, channelID, msg.Content, msg.ReplyToMessageID)
|
msgID, err := c.sendChunk(ctx, channelID, msg.Content, msg.ReplyToMessageID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return []string{msgID}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendMedia implements the channels.MediaSender interface.
|
// SendMedia implements the channels.MediaSender interface.
|
||||||
func (c *DiscordChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *DiscordChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
channelID := msg.ChatID
|
channelID := msg.ChatID
|
||||||
if channelID == "" {
|
if channelID == "" {
|
||||||
return fmt.Errorf("channel ID is empty")
|
return nil, fmt.Errorf("channel ID is empty")
|
||||||
}
|
}
|
||||||
|
|
||||||
store := c.GetMediaStore()
|
store := c.GetMediaStore()
|
||||||
if store == nil {
|
if store == nil {
|
||||||
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Collect all files into a single ChannelMessageSendComplex call
|
// Collect all files into a single ChannelMessageSendComplex call
|
||||||
|
|
@ -202,33 +206,41 @@ func (c *DiscordChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMes
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(files) == 0 {
|
if len(files) == 0 {
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
sendCtx, cancel := context.WithTimeout(ctx, sendTimeout)
|
sendCtx, cancel := context.WithTimeout(ctx, sendTimeout)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
done := make(chan error, 1)
|
type mediaResult struct {
|
||||||
|
id string
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
done := make(chan mediaResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
_, err := c.session.ChannelMessageSendComplex(channelID, &discordgo.MessageSend{
|
sentMsg, err := c.session.ChannelMessageSendComplex(channelID, &discordgo.MessageSend{
|
||||||
Content: caption,
|
Content: caption,
|
||||||
Files: files,
|
Files: files,
|
||||||
})
|
})
|
||||||
done <- err
|
if err != nil {
|
||||||
|
done <- mediaResult{err: err}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
done <- mediaResult{id: sentMsg.ID}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case err := <-done:
|
case r := <-done:
|
||||||
// Close all file readers
|
// Close all file readers
|
||||||
for _, f := range files {
|
for _, f := range files {
|
||||||
if closer, ok := f.Reader.(*os.File); ok {
|
if closer, ok := f.Reader.(*os.File); ok {
|
||||||
closer.Close()
|
closer.Close()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err != nil {
|
if r.err != nil {
|
||||||
return fmt.Errorf("discord send media: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("discord send media: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
return nil
|
return []string{r.id}, nil
|
||||||
case <-sendCtx.Done():
|
case <-sendCtx.Done():
|
||||||
// Close all file readers
|
// Close all file readers
|
||||||
for _, f := range files {
|
for _, f := range files {
|
||||||
|
|
@ -236,7 +248,7 @@ func (c *DiscordChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMes
|
||||||
closer.Close()
|
closer.Close()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return sendCtx.Err()
|
return nil, sendCtx.Err()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -264,18 +276,25 @@ func (c *DiscordChannel) SendPlaceholder(ctx context.Context, chatID string) (st
|
||||||
return msg.ID, nil
|
return msg.ID, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content, replyToID string) error {
|
func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content, replyToID string) (string, error) {
|
||||||
// Use the passed ctx for timeout control
|
// Use the passed ctx for timeout control
|
||||||
sendCtx, cancel := context.WithTimeout(ctx, sendTimeout)
|
sendCtx, cancel := context.WithTimeout(ctx, sendTimeout)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
done := make(chan error, 1)
|
type result struct {
|
||||||
|
id string
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
done := make(chan result, 1)
|
||||||
go func() {
|
go func() {
|
||||||
var err error
|
var (
|
||||||
|
msg *discordgo.Message
|
||||||
|
err error
|
||||||
|
)
|
||||||
|
|
||||||
// If we have an ID, we send the message as "Reply"
|
// If we have an ID, we send the message as "Reply"
|
||||||
if replyToID != "" {
|
if replyToID != "" {
|
||||||
_, err = c.session.ChannelMessageSendComplex(channelID, &discordgo.MessageSend{
|
msg, err = c.session.ChannelMessageSendComplex(channelID, &discordgo.MessageSend{
|
||||||
Content: content,
|
Content: content,
|
||||||
Reference: &discordgo.MessageReference{
|
Reference: &discordgo.MessageReference{
|
||||||
MessageID: replyToID,
|
MessageID: replyToID,
|
||||||
|
|
@ -284,20 +303,21 @@ func (c *DiscordChannel) sendChunk(ctx context.Context, channelID, content, repl
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
// Otherwise, we send a normal message
|
// Otherwise, we send a normal message
|
||||||
_, err = c.session.ChannelMessageSend(channelID, content)
|
msg, err = c.session.ChannelMessageSend(channelID, content)
|
||||||
}
|
}
|
||||||
|
|
||||||
done <- err
|
if err != nil {
|
||||||
|
done <- result{err: fmt.Errorf("discord send: %w", channels.ErrTemporary)}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
done <- result{id: msg.ID}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case err := <-done:
|
case r := <-done:
|
||||||
if err != nil {
|
return r.id, r.err
|
||||||
return fmt.Errorf("discord send: %w", channels.ErrTemporary)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
case <-sendCtx.Done():
|
case <-sendCtx.Done():
|
||||||
return sendCtx.Err()
|
return "", sendCtx.Err()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -36,8 +36,8 @@ func (c *FeishuChannel) Stop(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Send is a stub method to satisfy the Channel interface
|
// Send is a stub method to satisfy the Channel interface
|
||||||
func (c *FeishuChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *FeishuChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
return errUnsupported
|
return nil, errUnsupported
|
||||||
}
|
}
|
||||||
|
|
||||||
// EditMessage is a stub method to satisfy MessageEditor
|
// EditMessage is a stub method to satisfy MessageEditor
|
||||||
|
|
@ -56,6 +56,6 @@ func (c *FeishuChannel) ReactToMessage(ctx context.Context, chatID, messageID st
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendMedia is a stub method to satisfy MediaSender
|
// SendMedia is a stub method to satisfy MediaSender
|
||||||
func (c *FeishuChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *FeishuChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
return errUnsupported
|
return nil, errUnsupported
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -131,26 +131,26 @@ func (c *FeishuChannel) Stop(ctx context.Context) error {
|
||||||
|
|
||||||
// Send sends a message using Interactive Card format for markdown rendering.
|
// Send sends a message using Interactive Card format for markdown rendering.
|
||||||
// Falls back to plain text message if card sending fails (e.g., table limit exceeded).
|
// Falls back to plain text message if card sending fails (e.g., table limit exceeded).
|
||||||
func (c *FeishuChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *FeishuChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
if msg.ChatID == "" {
|
if msg.ChatID == "" {
|
||||||
return fmt.Errorf("chat ID is empty: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("chat ID is empty: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build interactive card with markdown content
|
// Build interactive card with markdown content
|
||||||
cardContent, err := buildMarkdownCard(msg.Content)
|
cardContent, err := buildMarkdownCard(msg.Content)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// If card build fails, fall back to plain text
|
// If card build fails, fall back to plain text
|
||||||
return c.sendText(ctx, msg.ChatID, msg.Content)
|
return nil, c.sendText(ctx, msg.ChatID, msg.Content)
|
||||||
}
|
}
|
||||||
|
|
||||||
// First attempt: try sending as interactive card
|
// First attempt: try sending as interactive card
|
||||||
err = c.sendCard(ctx, msg.ChatID, cardContent)
|
err = c.sendCard(ctx, msg.ChatID, cardContent)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check if error is due to card table limit (error code 11310)
|
// Check if error is due to card table limit (error code 11310)
|
||||||
|
|
@ -167,14 +167,14 @@ func (c *FeishuChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
// Second attempt: fall back to plain text message
|
// Second attempt: fall back to plain text message
|
||||||
textErr := c.sendText(ctx, msg.ChatID, msg.Content)
|
textErr := c.sendText(ctx, msg.ChatID, msg.Content)
|
||||||
if textErr == nil {
|
if textErr == nil {
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
// If text also fails, return the text error
|
// If text also fails, return the text error
|
||||||
return textErr
|
return nil, textErr
|
||||||
}
|
}
|
||||||
|
|
||||||
// For other errors, return the original card error
|
// For other errors, return the original card error
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// EditMessage implements channels.MessageEditor.
|
// EditMessage implements channels.MessageEditor.
|
||||||
|
|
@ -310,27 +310,27 @@ func (c *FeishuChannel) ReactToMessage(ctx context.Context, chatID, messageID st
|
||||||
|
|
||||||
// SendMedia implements channels.MediaSender.
|
// SendMedia implements channels.MediaSender.
|
||||||
// Uploads images/files via Feishu API then sends as messages.
|
// Uploads images/files via Feishu API then sends as messages.
|
||||||
func (c *FeishuChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *FeishuChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
if msg.ChatID == "" {
|
if msg.ChatID == "" {
|
||||||
return fmt.Errorf("chat ID is empty: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("chat ID is empty: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
store := c.GetMediaStore()
|
store := c.GetMediaStore()
|
||||||
if store == nil {
|
if store == nil {
|
||||||
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, part := range msg.Parts {
|
for _, part := range msg.Parts {
|
||||||
if err := c.sendMediaPart(ctx, msg.ChatID, part, store); err != nil {
|
if err := c.sendMediaPart(ctx, msg.ChatID, part, store); err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// sendMediaPart resolves and sends a single media part.
|
// sendMediaPart resolves and sends a single media part.
|
||||||
|
|
|
||||||
|
|
@ -130,18 +130,18 @@ func (c *IRCChannel) Stop(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Send sends a message to an IRC channel or user.
|
// Send sends a message to an IRC channel or user.
|
||||||
func (c *IRCChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *IRCChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
target := msg.ChatID
|
target := msg.ChatID
|
||||||
if target == "" {
|
if target == "" {
|
||||||
return fmt.Errorf("chat ID is empty: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("chat ID is empty: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
if strings.TrimSpace(msg.Content) == "" {
|
if strings.TrimSpace(msg.Content) == "" {
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Send each line separately (IRC is line-oriented)
|
// Send each line separately (IRC is line-oriented)
|
||||||
|
|
@ -158,7 +158,7 @@ func (c *IRCChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
"target": target,
|
"target": target,
|
||||||
"lines": len(lines),
|
"lines": len(lines),
|
||||||
})
|
})
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// StartTyping implements channels.TypingCapable using IRCv3 +typing client tag.
|
// StartTyping implements channels.TypingCapable using IRCv3 +typing client tag.
|
||||||
|
|
|
||||||
|
|
@ -496,9 +496,9 @@ func (c *LINEChannel) resolveChatID(source lineSource) string {
|
||||||
|
|
||||||
// Send sends a message to LINE. It first tries the Reply API (free)
|
// Send sends a message to LINE. It first tries the Reply API (free)
|
||||||
// using a cached reply token, then falls back to the Push API.
|
// using a cached reply token, then falls back to the Push API.
|
||||||
func (c *LINEChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *LINEChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
// Load and consume quote token for this chat
|
// Load and consume quote token for this chat
|
||||||
|
|
@ -516,28 +516,28 @@ func (c *LINEChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
"chat_id": msg.ChatID,
|
"chat_id": msg.ChatID,
|
||||||
"quoted": quoteToken != "",
|
"quoted": quoteToken != "",
|
||||||
})
|
})
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
logger.DebugC("line", "Reply API failed, falling back to Push API")
|
logger.DebugC("line", "Reply API failed, falling back to Push API")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Fall back to Push API
|
// Fall back to Push API
|
||||||
return c.sendPush(ctx, msg.ChatID, msg.Content, quoteToken)
|
return nil, c.sendPush(ctx, msg.ChatID, msg.Content, quoteToken)
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendMedia implements the channels.MediaSender interface.
|
// SendMedia implements the channels.MediaSender interface.
|
||||||
// LINE requires media to be accessible via public URL; since we only have local files,
|
// LINE requires media to be accessible via public URL; since we only have local files,
|
||||||
// we fall back to sending a text message with the filename/caption.
|
// we fall back to sending a text message with the filename/caption.
|
||||||
// For full support, an external file hosting service would be needed.
|
// For full support, an external file hosting service would be needed.
|
||||||
func (c *LINEChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *LINEChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
store := c.GetMediaStore()
|
store := c.GetMediaStore()
|
||||||
if store == nil {
|
if store == nil {
|
||||||
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
// LINE Messaging API requires publicly accessible URLs for media messages.
|
// LINE Messaging API requires publicly accessible URLs for media messages.
|
||||||
|
|
@ -549,11 +549,11 @@ func (c *LINEChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessag
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.sendPush(ctx, msg.ChatID, caption, ""); err != nil {
|
if err := c.sendPush(ctx, msg.ChatID, caption, ""); err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// buildTextMessage creates a text message object, optionally with quoteToken.
|
// buildTextMessage creates a text message object, optionally with quoteToken.
|
||||||
|
|
|
||||||
|
|
@ -240,15 +240,15 @@ func (c *MaixCamChannel) Stop(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *MaixCamChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *MaixCamChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check ctx before entering write path
|
// Check ctx before entering write path
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return ctx.Err()
|
return nil, ctx.Err()
|
||||||
default:
|
default:
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -257,7 +257,7 @@ func (c *MaixCamChannel) Send(ctx context.Context, msg bus.OutboundMessage) erro
|
||||||
|
|
||||||
if len(c.clients) == 0 {
|
if len(c.clients) == 0 {
|
||||||
logger.WarnC("maixcam", "No MaixCam devices connected")
|
logger.WarnC("maixcam", "No MaixCam devices connected")
|
||||||
return fmt.Errorf("no connected MaixCam devices")
|
return nil, fmt.Errorf("no connected MaixCam devices")
|
||||||
}
|
}
|
||||||
|
|
||||||
response := map[string]any{
|
response := map[string]any{
|
||||||
|
|
@ -269,7 +269,7 @@ func (c *MaixCamChannel) Send(ctx context.Context, msg bus.OutboundMessage) erro
|
||||||
|
|
||||||
data, err := json.Marshal(response)
|
data, err := json.Marshal(response)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to marshal response: %w", err)
|
return nil, fmt.Errorf("failed to marshal response: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
var sendErr error
|
var sendErr error
|
||||||
|
|
@ -285,5 +285,5 @@ func (c *MaixCamChannel) Send(ctx context.Context, msg bus.OutboundMessage) erro
|
||||||
_ = conn.SetWriteDeadline(time.Time{})
|
_ = conn.SetWriteDeadline(time.Time{})
|
||||||
}
|
}
|
||||||
|
|
||||||
return sendErr
|
return nil, sendErr
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -158,8 +158,8 @@ func (m *Manager) RecordReactionUndo(channel, chatID string, undo func()) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// preSend handles typing stop, reaction undo, and placeholder editing before sending a message.
|
// preSend handles typing stop, reaction undo, and placeholder editing before sending a message.
|
||||||
// Returns true if the message was already delivered (skip Send).
|
// Returns the delivered message IDs and true when delivery completed before a normal Send.
|
||||||
func (m *Manager) preSend(ctx context.Context, name string, msg bus.OutboundMessage, ch Channel) bool {
|
func (m *Manager) preSend(ctx context.Context, name string, msg bus.OutboundMessage, ch Channel) ([]string, bool) {
|
||||||
key := name + ":" + msg.ChatID
|
key := name + ":" + msg.ChatID
|
||||||
|
|
||||||
// 1. Stop typing
|
// 1. Stop typing
|
||||||
|
|
@ -188,7 +188,7 @@ func (m *Manager) preSend(ctx context.Context, name string, msg bus.OutboundMess
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return true
|
return nil, true
|
||||||
}
|
}
|
||||||
|
|
||||||
// 4. Try editing placeholder
|
// 4. Try editing placeholder
|
||||||
|
|
@ -196,14 +196,14 @@ func (m *Manager) preSend(ctx context.Context, name string, msg bus.OutboundMess
|
||||||
if entry, ok := v.(placeholderEntry); ok && entry.id != "" {
|
if entry, ok := v.(placeholderEntry); ok && entry.id != "" {
|
||||||
if editor, ok := ch.(MessageEditor); ok {
|
if editor, ok := ch.(MessageEditor); ok {
|
||||||
if err := editor.EditMessage(ctx, msg.ChatID, entry.id, msg.Content); err == nil {
|
if err := editor.EditMessage(ctx, msg.ChatID, entry.id, msg.Content); err == nil {
|
||||||
return true // edited successfully, skip Send
|
return []string{entry.id}, true
|
||||||
}
|
}
|
||||||
// edit failed → fall through to normal Send
|
// edit failed → fall through to normal Send
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return false
|
return nil, false
|
||||||
}
|
}
|
||||||
|
|
||||||
// preSendMedia handles typing stop, reaction undo, and placeholder cleanup
|
// preSendMedia handles typing stop, reaction undo, and placeholder cleanup
|
||||||
|
|
@ -699,23 +699,29 @@ func splitByLength(content string, maxLen int) []string {
|
||||||
// - ErrNotRunning / ErrSendFailed: permanent, no retry
|
// - ErrNotRunning / ErrSendFailed: permanent, no retry
|
||||||
// - ErrRateLimit: fixed delay retry
|
// - ErrRateLimit: fixed delay retry
|
||||||
// - ErrTemporary / unknown: exponential backoff retry
|
// - ErrTemporary / unknown: exponential backoff retry
|
||||||
func (m *Manager) sendWithRetry(ctx context.Context, name string, w *channelWorker, msg bus.OutboundMessage) {
|
func (m *Manager) sendWithRetry(
|
||||||
|
ctx context.Context,
|
||||||
|
name string,
|
||||||
|
w *channelWorker,
|
||||||
|
msg bus.OutboundMessage,
|
||||||
|
) ([]string, bool) {
|
||||||
// Rate limit: wait for token
|
// Rate limit: wait for token
|
||||||
if err := w.limiter.Wait(ctx); err != nil {
|
if err := w.limiter.Wait(ctx); err != nil {
|
||||||
// ctx canceled, shutting down
|
// ctx canceled, shutting down
|
||||||
return
|
return nil, false
|
||||||
}
|
}
|
||||||
|
|
||||||
// Pre-send: stop typing and try to edit placeholder
|
// Pre-send: stop typing and try to edit placeholder
|
||||||
if m.preSend(ctx, name, msg, w.ch) {
|
if msgIDs, handled := m.preSend(ctx, name, msg, w.ch); handled {
|
||||||
return // placeholder was edited successfully, skip Send
|
return msgIDs, true
|
||||||
}
|
}
|
||||||
|
|
||||||
var lastErr error
|
var lastErr error
|
||||||
|
var msgIDs []string
|
||||||
for attempt := 0; attempt <= maxRetries; attempt++ {
|
for attempt := 0; attempt <= maxRetries; attempt++ {
|
||||||
lastErr = w.ch.Send(ctx, msg)
|
msgIDs, lastErr = w.ch.Send(ctx, msg)
|
||||||
if lastErr == nil {
|
if lastErr == nil {
|
||||||
return
|
return msgIDs, true
|
||||||
}
|
}
|
||||||
|
|
||||||
// Permanent failures — don't retry
|
// Permanent failures — don't retry
|
||||||
|
|
@ -734,7 +740,7 @@ func (m *Manager) sendWithRetry(ctx context.Context, name string, w *channelWork
|
||||||
case <-time.After(rateLimitDelay):
|
case <-time.After(rateLimitDelay):
|
||||||
continue
|
continue
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return nil, false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -743,7 +749,7 @@ func (m *Manager) sendWithRetry(ctx context.Context, name string, w *channelWork
|
||||||
select {
|
select {
|
||||||
case <-time.After(backoff):
|
case <-time.After(backoff):
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return nil, false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -754,6 +760,8 @@ func (m *Manager) sendWithRetry(ctx context.Context, name string, w *channelWork
|
||||||
"error": lastErr.Error(),
|
"error": lastErr.Error(),
|
||||||
"retries": maxRetries,
|
"retries": maxRetries,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
return nil, false
|
||||||
}
|
}
|
||||||
|
|
||||||
func dispatchLoop[M any](
|
func dispatchLoop[M any](
|
||||||
|
|
@ -855,7 +863,7 @@ func (m *Manager) runMediaWorker(ctx context.Context, name string, w *channelWor
|
||||||
if !ok {
|
if !ok {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
_ = m.sendMediaWithRetry(ctx, name, w, msg)
|
_, _ = m.sendMediaWithRetry(ctx, name, w, msg)
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
@ -863,14 +871,14 @@ func (m *Manager) runMediaWorker(ctx context.Context, name string, w *channelWor
|
||||||
}
|
}
|
||||||
|
|
||||||
// sendMediaWithRetry sends a media message through the channel with rate limiting and
|
// sendMediaWithRetry sends a media message through the channel with rate limiting and
|
||||||
// retry logic. It returns nil on success, or the last error after retries,
|
// retry logic. It returns the message IDs and nil on success, or nil and the last error
|
||||||
// including when the channel does not support MediaSender.
|
// after retries, including when the channel does not support MediaSender.
|
||||||
func (m *Manager) sendMediaWithRetry(
|
func (m *Manager) sendMediaWithRetry(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
name string,
|
name string,
|
||||||
w *channelWorker,
|
w *channelWorker,
|
||||||
msg bus.OutboundMediaMessage,
|
msg bus.OutboundMediaMessage,
|
||||||
) error {
|
) ([]string, error) {
|
||||||
ms, ok := w.ch.(MediaSender)
|
ms, ok := w.ch.(MediaSender)
|
||||||
if !ok {
|
if !ok {
|
||||||
err := fmt.Errorf("channel %q does not support media sending", name)
|
err := fmt.Errorf("channel %q does not support media sending", name)
|
||||||
|
|
@ -878,22 +886,23 @@ func (m *Manager) sendMediaWithRetry(
|
||||||
"channel": name,
|
"channel": name,
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Rate limit: wait for token
|
// Rate limit: wait for token
|
||||||
if err := w.limiter.Wait(ctx); err != nil {
|
if err := w.limiter.Wait(ctx); err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Pre-send: stop typing and clean up any placeholder before sending media.
|
// Pre-send: stop typing and clean up any placeholder before sending media.
|
||||||
m.preSendMedia(ctx, name, msg, w.ch)
|
m.preSendMedia(ctx, name, msg, w.ch)
|
||||||
|
|
||||||
var lastErr error
|
var lastErr error
|
||||||
|
var msgIDs []string
|
||||||
for attempt := 0; attempt <= maxRetries; attempt++ {
|
for attempt := 0; attempt <= maxRetries; attempt++ {
|
||||||
lastErr = ms.SendMedia(ctx, msg)
|
msgIDs, lastErr = ms.SendMedia(ctx, msg)
|
||||||
if lastErr == nil {
|
if lastErr == nil {
|
||||||
return nil
|
return msgIDs, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Permanent failures — don't retry
|
// Permanent failures — don't retry
|
||||||
|
|
@ -912,7 +921,7 @@ func (m *Manager) sendMediaWithRetry(
|
||||||
case <-time.After(rateLimitDelay):
|
case <-time.After(rateLimitDelay):
|
||||||
continue
|
continue
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return ctx.Err()
|
return nil, ctx.Err()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -921,7 +930,7 @@ func (m *Manager) sendMediaWithRetry(
|
||||||
select {
|
select {
|
||||||
case <-time.After(backoff):
|
case <-time.After(backoff):
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return ctx.Err()
|
return nil, ctx.Err()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -932,7 +941,7 @@ func (m *Manager) sendMediaWithRetry(
|
||||||
"error": lastErr.Error(),
|
"error": lastErr.Error(),
|
||||||
"retries": maxRetries,
|
"retries": maxRetries,
|
||||||
})
|
})
|
||||||
return lastErr
|
return nil, lastErr
|
||||||
}
|
}
|
||||||
|
|
||||||
// runTTLJanitor periodically scans the typingStops and placeholders maps
|
// runTTLJanitor periodically scans the typingStops and placeholders maps
|
||||||
|
|
@ -1166,7 +1175,8 @@ func (m *Manager) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) e
|
||||||
return fmt.Errorf("channel %s has no active worker", msg.Channel)
|
return fmt.Errorf("channel %s has no active worker", msg.Channel)
|
||||||
}
|
}
|
||||||
|
|
||||||
return m.sendMediaWithRetry(ctx, msg.Channel, w, msg)
|
_, err := m.sendMediaWithRetry(ctx, msg.Channel, w, msg)
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *Manager) SendToChannel(ctx context.Context, channelName, chatID, content string) error {
|
func (m *Manager) SendToChannel(ctx context.Context, channelName, chatID, content string) error {
|
||||||
|
|
@ -1196,5 +1206,6 @@ func (m *Manager) SendToChannel(ctx context.Context, channelName, chatID, conten
|
||||||
|
|
||||||
// Fallback: direct send (should not happen)
|
// Fallback: direct send (should not happen)
|
||||||
channel, _ := m.channels[channelName]
|
channel, _ := m.channels[channelName]
|
||||||
return channel.Send(ctx, msg)
|
_, err := channel.Send(ctx, msg)
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -25,9 +25,12 @@ type mockChannel struct {
|
||||||
lastPlaceholderID string
|
lastPlaceholderID string
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *mockChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (m *mockChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
m.sentMessages = append(m.sentMessages, msg)
|
m.sentMessages = append(m.sentMessages, msg)
|
||||||
return m.sendFn(ctx, msg)
|
if m.sendFn == nil {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return nil, m.sendFn(ctx, msg)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *mockChannel) Start(ctx context.Context) error { return nil }
|
func (m *mockChannel) Start(ctx context.Context) error { return nil }
|
||||||
|
|
@ -46,16 +49,16 @@ func (m *mockChannel) EditMessage(ctx context.Context, chatID, messageID, conten
|
||||||
|
|
||||||
type mockMediaChannel struct {
|
type mockMediaChannel struct {
|
||||||
mockChannel
|
mockChannel
|
||||||
sendMediaFn func(ctx context.Context, msg bus.OutboundMediaMessage) error
|
sendMediaFn func(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error)
|
||||||
sentMediaMessages []bus.OutboundMediaMessage
|
sentMediaMessages []bus.OutboundMediaMessage
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *mockMediaChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (m *mockMediaChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
m.sentMediaMessages = append(m.sentMediaMessages, msg)
|
m.sentMediaMessages = append(m.sentMediaMessages, msg)
|
||||||
if m.sendMediaFn != nil {
|
if m.sendMediaFn != nil {
|
||||||
return m.sendMediaFn(ctx, msg)
|
return m.sendMediaFn(ctx, msg)
|
||||||
}
|
}
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type mockDeletingMediaChannel struct {
|
type mockDeletingMediaChannel struct {
|
||||||
|
|
@ -247,9 +250,9 @@ func TestSendMedia_Success(t *testing.T) {
|
||||||
m := newTestManager()
|
m := newTestManager()
|
||||||
var callCount int
|
var callCount int
|
||||||
ch := &mockMediaChannel{
|
ch := &mockMediaChannel{
|
||||||
sendMediaFn: func(_ context.Context, _ bus.OutboundMediaMessage) error {
|
sendMediaFn: func(_ context.Context, _ bus.OutboundMediaMessage) ([]string, error) {
|
||||||
callCount++
|
callCount++
|
||||||
return nil
|
return nil, nil
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
w := &channelWorker{
|
w := &channelWorker{
|
||||||
|
|
@ -275,8 +278,8 @@ func TestSendMedia_Success(t *testing.T) {
|
||||||
func TestSendMedia_PropagatesFailure(t *testing.T) {
|
func TestSendMedia_PropagatesFailure(t *testing.T) {
|
||||||
m := newTestManager()
|
m := newTestManager()
|
||||||
ch := &mockMediaChannel{
|
ch := &mockMediaChannel{
|
||||||
sendMediaFn: func(_ context.Context, _ bus.OutboundMediaMessage) error {
|
sendMediaFn: func(_ context.Context, _ bus.OutboundMediaMessage) ([]string, error) {
|
||||||
return fmt.Errorf("bad upload: %w", ErrSendFailed)
|
return nil, fmt.Errorf("bad upload: %w", ErrSendFailed)
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
w := &channelWorker{
|
w := &channelWorker{
|
||||||
|
|
@ -330,8 +333,8 @@ func TestSendMedia_DeletesPlaceholderBeforeSending(t *testing.T) {
|
||||||
m := newTestManager()
|
m := newTestManager()
|
||||||
ch := &mockDeletingMediaChannel{
|
ch := &mockDeletingMediaChannel{
|
||||||
mockMediaChannel: mockMediaChannel{
|
mockMediaChannel: mockMediaChannel{
|
||||||
sendMediaFn: func(_ context.Context, _ bus.OutboundMediaMessage) error {
|
sendMediaFn: func(_ context.Context, _ bus.OutboundMediaMessage) ([]string, error) {
|
||||||
return nil
|
return nil, nil
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
@ -628,7 +631,7 @@ func TestPreSend_PlaceholderEditSuccess(t *testing.T) {
|
||||||
m.RecordPlaceholder("test", "123", "456")
|
m.RecordPlaceholder("test", "123", "456")
|
||||||
|
|
||||||
msg := bus.OutboundMessage{Channel: "test", ChatID: "123", Content: "hello"}
|
msg := bus.OutboundMessage{Channel: "test", ChatID: "123", Content: "hello"}
|
||||||
edited := m.preSend(context.Background(), "test", msg, ch)
|
_, edited := m.preSend(context.Background(), "test", msg, ch)
|
||||||
|
|
||||||
if !edited {
|
if !edited {
|
||||||
t.Fatal("expected preSend to return true (placeholder edited)")
|
t.Fatal("expected preSend to return true (placeholder edited)")
|
||||||
|
|
@ -658,7 +661,7 @@ func TestPreSend_PlaceholderEditFails_FallsThrough(t *testing.T) {
|
||||||
m.RecordPlaceholder("test", "123", "456")
|
m.RecordPlaceholder("test", "123", "456")
|
||||||
|
|
||||||
msg := bus.OutboundMessage{Channel: "test", ChatID: "123", Content: "hello"}
|
msg := bus.OutboundMessage{Channel: "test", ChatID: "123", Content: "hello"}
|
||||||
edited := m.preSend(context.Background(), "test", msg, ch)
|
_, edited := m.preSend(context.Background(), "test", msg, ch)
|
||||||
|
|
||||||
if edited {
|
if edited {
|
||||||
t.Fatal("expected preSend to return false when edit fails")
|
t.Fatal("expected preSend to return false when edit fails")
|
||||||
|
|
@ -734,7 +737,7 @@ func TestPreSend_NoRegisteredState(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
msg := bus.OutboundMessage{Channel: "test", ChatID: "123", Content: "hello"}
|
msg := bus.OutboundMessage{Channel: "test", ChatID: "123", Content: "hello"}
|
||||||
edited := m.preSend(context.Background(), "test", msg, ch)
|
_, edited := m.preSend(context.Background(), "test", msg, ch)
|
||||||
|
|
||||||
if edited {
|
if edited {
|
||||||
t.Fatal("expected preSend to return false with no registered state")
|
t.Fatal("expected preSend to return false with no registered state")
|
||||||
|
|
@ -764,7 +767,7 @@ func TestPreSend_TypingAndPlaceholder(t *testing.T) {
|
||||||
m.RecordPlaceholder("test", "123", "456")
|
m.RecordPlaceholder("test", "123", "456")
|
||||||
|
|
||||||
msg := bus.OutboundMessage{Channel: "test", ChatID: "123", Content: "hello"}
|
msg := bus.OutboundMessage{Channel: "test", ChatID: "123", Content: "hello"}
|
||||||
edited := m.preSend(context.Background(), "test", msg, ch)
|
_, edited := m.preSend(context.Background(), "test", msg, ch)
|
||||||
|
|
||||||
if !stopCalled {
|
if !stopCalled {
|
||||||
t.Fatal("expected typing stop to be called")
|
t.Fatal("expected typing stop to be called")
|
||||||
|
|
@ -1025,7 +1028,7 @@ func TestPreSendStillWorksWithWrappedTypes(t *testing.T) {
|
||||||
m.RecordPlaceholder("test", "chat1", "ph_id")
|
m.RecordPlaceholder("test", "chat1", "ph_id")
|
||||||
|
|
||||||
msg := bus.OutboundMessage{Channel: "test", ChatID: "chat1", Content: "response"}
|
msg := bus.OutboundMessage{Channel: "test", ChatID: "chat1", Content: "response"}
|
||||||
edited := m.preSend(context.Background(), "test", msg, ch)
|
_, edited := m.preSend(context.Background(), "test", msg, ch)
|
||||||
|
|
||||||
if !stopCalled {
|
if !stopCalled {
|
||||||
t.Fatal("expected typing stop to be called via wrapped type")
|
t.Fatal("expected typing stop to be called via wrapped type")
|
||||||
|
|
|
||||||
|
|
@ -380,26 +380,26 @@ func markdownToHTML(md string) string {
|
||||||
return strings.TrimSpace(string(markdown.ToHTML([]byte(md), p, renderer)))
|
return strings.TrimSpace(string(markdown.ToHTML([]byte(md), p, renderer)))
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *MatrixChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *MatrixChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
roomID := id.RoomID(strings.TrimSpace(msg.ChatID))
|
roomID := id.RoomID(strings.TrimSpace(msg.ChatID))
|
||||||
if roomID == "" {
|
if roomID == "" {
|
||||||
return fmt.Errorf("matrix room ID is empty: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("matrix room ID is empty: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
content := strings.TrimSpace(msg.Content)
|
content := strings.TrimSpace(msg.Content)
|
||||||
if content == "" {
|
if content == "" {
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err := c.client.SendMessageEvent(ctx, roomID, event.EventMessage, c.messageContent(content))
|
resp, err := c.client.SendMessageEvent(ctx, roomID, event.EventMessage, c.messageContent(content))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("matrix send: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("matrix send: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
return nil
|
return []string{resp.EventID.String()}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *MatrixChannel) messageContent(text string) *event.MessageEventContent {
|
func (c *MatrixChannel) messageContent(text string) *event.MessageEventContent {
|
||||||
|
|
@ -412,9 +412,9 @@ func (c *MatrixChannel) messageContent(text string) *event.MessageEventContent {
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendMedia implements channels.MediaSender.
|
// SendMedia implements channels.MediaSender.
|
||||||
func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
sendCtx := ctx
|
sendCtx := ctx
|
||||||
if sendCtx == nil {
|
if sendCtx == nil {
|
||||||
|
|
@ -423,17 +423,18 @@ func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
|
|
||||||
roomID := id.RoomID(strings.TrimSpace(msg.ChatID))
|
roomID := id.RoomID(strings.TrimSpace(msg.ChatID))
|
||||||
if roomID == "" {
|
if roomID == "" {
|
||||||
return fmt.Errorf("matrix room ID is empty: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("matrix room ID is empty: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
store := c.GetMediaStore()
|
store := c.GetMediaStore()
|
||||||
if store == nil {
|
if store == nil {
|
||||||
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var eventIDs []string
|
||||||
for _, part := range msg.Parts {
|
for _, part := range msg.Parts {
|
||||||
if err := sendCtx.Err(); err != nil {
|
if err := sendCtx.Err(); err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
localPath, meta, err := store.ResolveWithMeta(part.Ref)
|
localPath, meta, err := store.ResolveWithMeta(part.Ref)
|
||||||
|
|
@ -498,7 +499,7 @@ func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
"type": part.Type,
|
"type": part.Type,
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return fmt.Errorf("matrix upload media: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("matrix upload media: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
||||||
msgType := matrixOutboundMsgType(part.Type, filename, contentType)
|
msgType := matrixOutboundMsgType(part.Type, filename, contentType)
|
||||||
|
|
@ -511,17 +512,21 @@ func (c *MatrixChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
uploadResp.ContentURI.CUString(),
|
uploadResp.ContentURI.CUString(),
|
||||||
)
|
)
|
||||||
|
|
||||||
if _, err := c.client.SendMessageEvent(sendCtx, roomID, event.EventMessage, content); err != nil {
|
sendResp, err := c.client.SendMessageEvent(sendCtx, roomID, event.EventMessage, content)
|
||||||
|
if err != nil {
|
||||||
logger.ErrorCF("matrix", "Failed to send media message", map[string]any{
|
logger.ErrorCF("matrix", "Failed to send media message", map[string]any{
|
||||||
"room_id": roomID.String(),
|
"room_id": roomID.String(),
|
||||||
"type": msgType,
|
"type": msgType,
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return fmt.Errorf("matrix send media: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("matrix send media: %w", channels.ErrTemporary)
|
||||||
|
}
|
||||||
|
if sendResp != nil {
|
||||||
|
eventIDs = append(eventIDs, sendResp.EventID.String())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return eventIDs, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// StartTyping implements channels.TypingCapable.
|
// StartTyping implements channels.TypingCapable.
|
||||||
|
|
|
||||||
|
|
@ -11,5 +11,5 @@ import (
|
||||||
// Manager discovers channels implementing this interface via type
|
// Manager discovers channels implementing this interface via type
|
||||||
// assertion and routes OutboundMediaMessage to them.
|
// assertion and routes OutboundMediaMessage to them.
|
||||||
type MediaSender interface {
|
type MediaSender interface {
|
||||||
SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error
|
SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -391,15 +391,15 @@ func (c *OneBotChannel) Stop(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *OneBotChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check ctx before entering write path
|
// Check ctx before entering write path
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return ctx.Err()
|
return nil, ctx.Err()
|
||||||
default:
|
default:
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -408,12 +408,12 @@ func (c *OneBotChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
c.mu.Unlock()
|
c.mu.Unlock()
|
||||||
|
|
||||||
if conn == nil {
|
if conn == nil {
|
||||||
return fmt.Errorf("OneBot WebSocket not connected")
|
return nil, fmt.Errorf("OneBot WebSocket not connected")
|
||||||
}
|
}
|
||||||
|
|
||||||
action, params, err := c.buildSendRequest(msg)
|
action, params, err := c.buildSendRequest(msg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
echo := fmt.Sprintf("send_%d", atomic.AddInt64(&c.echoCounter, 1))
|
echo := fmt.Sprintf("send_%d", atomic.AddInt64(&c.echoCounter, 1))
|
||||||
|
|
@ -426,7 +426,7 @@ func (c *OneBotChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
|
|
||||||
data, err := json.Marshal(req)
|
data, err := json.Marshal(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to marshal OneBot request: %w", err)
|
return nil, fmt.Errorf("failed to marshal OneBot request: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
c.writeMu.Lock()
|
c.writeMu.Lock()
|
||||||
|
|
@ -439,21 +439,21 @@ func (c *OneBotChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
logger.ErrorCF("onebot", "Failed to send message", map[string]any{
|
logger.ErrorCF("onebot", "Failed to send message", map[string]any{
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return fmt.Errorf("onebot send: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("onebot send: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendMedia implements the channels.MediaSender interface.
|
// SendMedia implements the channels.MediaSender interface.
|
||||||
func (c *OneBotChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *OneBotChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return ctx.Err()
|
return nil, ctx.Err()
|
||||||
default:
|
default:
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -462,12 +462,12 @@ func (c *OneBotChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
c.mu.Unlock()
|
c.mu.Unlock()
|
||||||
|
|
||||||
if conn == nil {
|
if conn == nil {
|
||||||
return fmt.Errorf("OneBot WebSocket not connected")
|
return nil, fmt.Errorf("OneBot WebSocket not connected")
|
||||||
}
|
}
|
||||||
|
|
||||||
store := c.GetMediaStore()
|
store := c.GetMediaStore()
|
||||||
if store == nil {
|
if store == nil {
|
||||||
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build media segments
|
// Build media segments
|
||||||
|
|
@ -508,7 +508,7 @@ func (c *OneBotChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(segments) == 0 {
|
if len(segments) == 0 {
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
chatID := msg.ChatID
|
chatID := msg.ChatID
|
||||||
|
|
@ -524,7 +524,7 @@ func (c *OneBotChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
|
|
||||||
id, err := strconv.ParseInt(rawID, 10, 64)
|
id, err := strconv.ParseInt(rawID, 10, 64)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid %s in chatID: %s: %w", idKey, chatID, channels.ErrSendFailed)
|
return nil, fmt.Errorf("invalid %s in chatID: %s: %w", idKey, chatID, channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
echo := fmt.Sprintf("send_%d", atomic.AddInt64(&c.echoCounter, 1))
|
echo := fmt.Sprintf("send_%d", atomic.AddInt64(&c.echoCounter, 1))
|
||||||
|
|
@ -537,7 +537,7 @@ func (c *OneBotChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
|
|
||||||
data, err := json.Marshal(req)
|
data, err := json.Marshal(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to marshal OneBot request: %w", err)
|
return nil, fmt.Errorf("failed to marshal OneBot request: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
c.writeMu.Lock()
|
c.writeMu.Lock()
|
||||||
|
|
@ -550,10 +550,10 @@ func (c *OneBotChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
logger.ErrorCF("onebot", "Failed to send media message", map[string]any{
|
logger.ErrorCF("onebot", "Failed to send media message", map[string]any{
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return fmt.Errorf("onebot send media: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("onebot send media: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) buildMessageSegments(chatID, content string) []oneBotMessageSegment {
|
func (c *OneBotChannel) buildMessageSegments(chatID, content string) []oneBotMessageSegment {
|
||||||
|
|
|
||||||
|
|
@ -273,22 +273,22 @@ func (c *PicoClientChannel) handleServerMessage(pc *picoConn, msg PicoMessage) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Send sends a message to the remote server.
|
// Send sends a message to the remote server.
|
||||||
func (c *PicoClientChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *PicoClientChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
pc := c.conn
|
pc := c.conn
|
||||||
c.mu.Unlock()
|
c.mu.Unlock()
|
||||||
if pc == nil || pc.closed.Load() {
|
if pc == nil || pc.closed.Load() {
|
||||||
return channels.ErrSendFailed
|
return nil, channels.ErrSendFailed
|
||||||
}
|
}
|
||||||
|
|
||||||
outMsg := newMessage(TypeMessageSend, map[string]any{
|
outMsg := newMessage(TypeMessageSend, map[string]any{
|
||||||
"content": msg.Content,
|
"content": msg.Content,
|
||||||
})
|
})
|
||||||
outMsg.SessionID = strings.TrimPrefix(msg.ChatID, "pico_client:")
|
outMsg.SessionID = strings.TrimPrefix(msg.ChatID, "pico_client:")
|
||||||
return pc.writeJSON(outMsg)
|
return nil, pc.writeJSON(outMsg)
|
||||||
}
|
}
|
||||||
|
|
||||||
// StartTyping implements channels.TypingCapable.
|
// StartTyping implements channels.TypingCapable.
|
||||||
|
|
|
||||||
|
|
@ -46,7 +46,7 @@ func TestSend_NotRunning(t *testing.T) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
err = ch.Send(context.Background(), bus.OutboundMessage{Content: "hi"})
|
_, err = ch.Send(context.Background(), bus.OutboundMessage{Content: "hi"})
|
||||||
if !errors.Is(err, channels.ErrNotRunning) {
|
if !errors.Is(err, channels.ErrNotRunning) {
|
||||||
t.Fatalf("expected ErrNotRunning, got %v", err)
|
t.Fatalf("expected ErrNotRunning, got %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -124,7 +124,7 @@ func TestClientChannel_ConnectAndSend(t *testing.T) {
|
||||||
defer ch.Stop(ctx)
|
defer ch.Stop(ctx)
|
||||||
|
|
||||||
// Send a message
|
// Send a message
|
||||||
err = ch.Send(ctx, bus.OutboundMessage{
|
_, err = ch.Send(ctx, bus.OutboundMessage{
|
||||||
ChatID: "pico_client:sess-1",
|
ChatID: "pico_client:sess-1",
|
||||||
Content: "hello",
|
Content: "hello",
|
||||||
})
|
})
|
||||||
|
|
@ -179,7 +179,7 @@ func TestClientChannel_ReceivesServerMessage(t *testing.T) {
|
||||||
defer ch.Stop(ctx)
|
defer ch.Stop(ctx)
|
||||||
|
|
||||||
// Send a message; the echo server replies with message.create
|
// Send a message; the echo server replies with message.create
|
||||||
err = ch.Send(ctx, bus.OutboundMessage{
|
_, err = ch.Send(ctx, bus.OutboundMessage{
|
||||||
ChatID: "pico_client:sess-echo",
|
ChatID: "pico_client:sess-echo",
|
||||||
Content: "ping",
|
Content: "ping",
|
||||||
})
|
})
|
||||||
|
|
@ -252,7 +252,7 @@ func TestSend_ClosedConnection(t *testing.T) {
|
||||||
ch.conn.close()
|
ch.conn.close()
|
||||||
ch.mu.Unlock()
|
ch.mu.Unlock()
|
||||||
|
|
||||||
err = ch.Send(ctx, bus.OutboundMessage{
|
_, err = ch.Send(ctx, bus.OutboundMessage{
|
||||||
ChatID: "pico_client:sess-close",
|
ChatID: "pico_client:sess-close",
|
||||||
Content: "should fail",
|
Content: "should fail",
|
||||||
})
|
})
|
||||||
|
|
|
||||||
|
|
@ -234,16 +234,16 @@ func (c *PicoChannel) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Send implements Channel — sends a message to the appropriate WebSocket connection.
|
// Send implements Channel — sends a message to the appropriate WebSocket connection.
|
||||||
func (c *PicoChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *PicoChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
outMsg := newMessage(TypeMessageCreate, map[string]any{
|
outMsg := newMessage(TypeMessageCreate, map[string]any{
|
||||||
"content": msg.Content,
|
"content": msg.Content,
|
||||||
})
|
})
|
||||||
|
|
||||||
return c.broadcastToSession(msg.ChatID, outMsg)
|
return nil, c.broadcastToSession(msg.ChatID, outMsg)
|
||||||
}
|
}
|
||||||
|
|
||||||
// EditMessage implements channels.MessageEditor.
|
// EditMessage implements channels.MessageEditor.
|
||||||
|
|
|
||||||
|
|
@ -200,9 +200,9 @@ func (c *QQChannel) getChatKind(chatID string) string {
|
||||||
return "group"
|
return "group"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *QQChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *QQChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
chatKind := c.getChatKind(msg.ChatID)
|
chatKind := c.getChatKind(msg.ChatID)
|
||||||
|
|
@ -236,11 +236,14 @@ func (c *QQChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Route to group or C2C.
|
// Route to group or C2C.
|
||||||
var err error
|
var (
|
||||||
|
sentMsg *dto.Message
|
||||||
|
err error
|
||||||
|
)
|
||||||
if chatKind == "group" {
|
if chatKind == "group" {
|
||||||
_, err = c.api.PostGroupMessage(ctx, msg.ChatID, msgToCreate)
|
sentMsg, err = c.api.PostGroupMessage(ctx, msg.ChatID, msgToCreate)
|
||||||
} else {
|
} else {
|
||||||
_, err = c.api.PostC2CMessage(ctx, msg.ChatID, msgToCreate)
|
sentMsg, err = c.api.PostC2CMessage(ctx, msg.ChatID, msgToCreate)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -249,10 +252,13 @@ func (c *QQChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
"chat_kind": chatKind,
|
"chat_kind": chatKind,
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return fmt.Errorf("qq send: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("qq send: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
if sentMsg == nil {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return []string{sentMsg.ID}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// StartTyping implements channels.TypingCapable.
|
// StartTyping implements channels.TypingCapable.
|
||||||
|
|
@ -319,13 +325,14 @@ func (c *QQChannel) StartTyping(ctx context.Context, chatID string) (func(), err
|
||||||
// QQ group/C2C media sending is a two-step flow:
|
// QQ group/C2C media sending is a two-step flow:
|
||||||
// 1. Upload media to /files using a remote URL or base64-encoded local bytes.
|
// 1. Upload media to /files using a remote URL or base64-encoded local bytes.
|
||||||
// 2. Send a msg_type=7 message using the returned file_info.
|
// 2. Send a msg_type=7 message using the returned file_info.
|
||||||
func (c *QQChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *QQChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
chatKind := c.getChatKind(msg.ChatID)
|
chatKind := c.getChatKind(msg.ChatID)
|
||||||
|
|
||||||
|
var messageIDs []string
|
||||||
for _, part := range msg.Parts {
|
for _, part := range msg.Parts {
|
||||||
fileInfo, err := c.uploadMedia(ctx, chatKind, msg.ChatID, part)
|
fileInfo, err := c.uploadMedia(ctx, chatKind, msg.ChatID, part)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -335,22 +342,26 @@ func (c *QQChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage)
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
if errors.Is(err, channels.ErrSendFailed) {
|
if errors.Is(err, channels.ErrSendFailed) {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
return fmt.Errorf("qq send media: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("qq send media: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.sendUploadedMedia(ctx, chatKind, msg.ChatID, part, fileInfo); err != nil {
|
sentMsg, err := c.sendUploadedMedia(ctx, chatKind, msg.ChatID, part, fileInfo)
|
||||||
|
if err != nil {
|
||||||
logger.ErrorCF("qq", "Failed to send media", map[string]any{
|
logger.ErrorCF("qq", "Failed to send media", map[string]any{
|
||||||
"type": part.Type,
|
"type": part.Type,
|
||||||
"chat_id": msg.ChatID,
|
"chat_id": msg.ChatID,
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return fmt.Errorf("qq send media: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("qq send media: %w", channels.ErrTemporary)
|
||||||
|
}
|
||||||
|
if sentMsg != nil && sentMsg.ID != "" {
|
||||||
|
messageIDs = append(messageIDs, sentMsg.ID)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return messageIDs, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type qqMediaUpload struct {
|
type qqMediaUpload struct {
|
||||||
|
|
@ -517,7 +528,7 @@ func (c *QQChannel) sendUploadedMedia(
|
||||||
chatKind, chatID string,
|
chatKind, chatID string,
|
||||||
part bus.MediaPart,
|
part bus.MediaPart,
|
||||||
fileInfo []byte,
|
fileInfo []byte,
|
||||||
) error {
|
) (*dto.Message, error) {
|
||||||
msg := &dto.MessageToCreate{
|
msg := &dto.MessageToCreate{
|
||||||
Content: part.Caption,
|
Content: part.Caption,
|
||||||
MsgType: dto.RichMediaMsg,
|
MsgType: dto.RichMediaMsg,
|
||||||
|
|
@ -532,11 +543,11 @@ func (c *QQChannel) sendUploadedMedia(
|
||||||
}
|
}
|
||||||
|
|
||||||
if chatKind == "group" {
|
if chatKind == "group" {
|
||||||
_, err := c.api.PostGroupMessage(ctx, chatID, msg)
|
sentMsg, err := c.api.PostGroupMessage(ctx, chatID, msg)
|
||||||
return err
|
return sentMsg, err
|
||||||
}
|
}
|
||||||
_, err := c.api.PostC2CMessage(ctx, chatID, msg)
|
sentMsg, err := c.api.PostC2CMessage(ctx, chatID, msg)
|
||||||
return err
|
return sentMsg, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *QQChannel) applyPassiveReplyMetadata(chatID string, msg *dto.MessageToCreate) {
|
func (c *QQChannel) applyPassiveReplyMetadata(chatID string, msg *dto.MessageToCreate) {
|
||||||
|
|
|
||||||
|
|
@ -209,7 +209,7 @@ func TestSendMedia_UploadsLocalFileAsBase64(t *testing.T) {
|
||||||
ch.lastMsgID.Store("group-1", "msg-1")
|
ch.lastMsgID.Store("group-1", "msg-1")
|
||||||
ch.msgSeqCounters.Store("group-1", new(atomic.Uint64))
|
ch.msgSeqCounters.Store("group-1", new(atomic.Uint64))
|
||||||
|
|
||||||
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
ChatID: "group-1",
|
ChatID: "group-1",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
Type: "image",
|
Type: "image",
|
||||||
|
|
@ -303,7 +303,7 @@ func assertAudioWAVUploadType(t *testing.T, duration time.Duration, wantFileType
|
||||||
ch.SetMediaStore(store)
|
ch.SetMediaStore(store)
|
||||||
ch.chatType.Store("group-1", "group")
|
ch.chatType.Store("group-1", "group")
|
||||||
|
|
||||||
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
ChatID: "group-1",
|
ChatID: "group-1",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
Type: "audio",
|
Type: "audio",
|
||||||
|
|
@ -337,7 +337,7 @@ func TestSendMedia_RemoteAudioFallsBackToFileUpload(t *testing.T) {
|
||||||
ch.SetRunning(true)
|
ch.SetRunning(true)
|
||||||
ch.chatType.Store("user-1", "direct")
|
ch.chatType.Store("user-1", "direct")
|
||||||
|
|
||||||
err := ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err := ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
ChatID: "user-1",
|
ChatID: "user-1",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
Type: "audio",
|
Type: "audio",
|
||||||
|
|
@ -383,7 +383,7 @@ func TestSendMedia_LocalAudioWithUnknownDurationFallsBackToFileUpload(t *testing
|
||||||
ch.SetMediaStore(store)
|
ch.SetMediaStore(store)
|
||||||
ch.chatType.Store("group-1", "group")
|
ch.chatType.Store("group-1", "group")
|
||||||
|
|
||||||
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
ChatID: "group-1",
|
ChatID: "group-1",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
Type: "audio",
|
Type: "audio",
|
||||||
|
|
@ -417,7 +417,7 @@ func TestSendMedia_UsesRemoteURLUploadForC2C(t *testing.T) {
|
||||||
ch.SetRunning(true)
|
ch.SetRunning(true)
|
||||||
ch.chatType.Store("user-1", "direct")
|
ch.chatType.Store("user-1", "direct")
|
||||||
|
|
||||||
err := ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err := ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
ChatID: "user-1",
|
ChatID: "user-1",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
Type: "file",
|
Type: "file",
|
||||||
|
|
@ -490,7 +490,7 @@ func TestSendMedia_LocalFileUploadIncludesStoredFilename(t *testing.T) {
|
||||||
ch.SetMediaStore(store)
|
ch.SetMediaStore(store)
|
||||||
ch.chatType.Store("user-1", "direct")
|
ch.chatType.Store("user-1", "direct")
|
||||||
|
|
||||||
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
ChatID: "user-1",
|
ChatID: "user-1",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
Type: "file",
|
Type: "file",
|
||||||
|
|
@ -528,7 +528,7 @@ func TestSendMedia_ReturnsSendFailedWithoutMediaStore(t *testing.T) {
|
||||||
ch.SetRunning(true)
|
ch.SetRunning(true)
|
||||||
ch.chatType.Store("group-1", "group")
|
ch.chatType.Store("group-1", "group")
|
||||||
|
|
||||||
err := ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err := ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
ChatID: "group-1",
|
ChatID: "group-1",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
Type: "image",
|
Type: "image",
|
||||||
|
|
@ -578,7 +578,7 @@ func TestSendMedia_ReturnsSendFailedWhenLocalFileExceedsBase64MiBLimit(t *testin
|
||||||
ch.SetMediaStore(store)
|
ch.SetMediaStore(store)
|
||||||
ch.chatType.Store("group-1", "group")
|
ch.chatType.Store("group-1", "group")
|
||||||
|
|
||||||
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
ChatID: "group-1",
|
ChatID: "group-1",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
Type: "file",
|
Type: "file",
|
||||||
|
|
|
||||||
|
|
@ -108,14 +108,14 @@ func (c *SlackChannel) Stop(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *SlackChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *SlackChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
channelID, threadTS := parseSlackChatID(msg.ChatID)
|
channelID, threadTS := parseSlackChatID(msg.ChatID)
|
||||||
if channelID == "" {
|
if channelID == "" {
|
||||||
return fmt.Errorf("invalid slack chat ID: %s", msg.ChatID)
|
return nil, fmt.Errorf("invalid slack chat ID: %s", msg.ChatID)
|
||||||
}
|
}
|
||||||
|
|
||||||
opts := []slack.MsgOption{
|
opts := []slack.MsgOption{
|
||||||
|
|
@ -130,9 +130,9 @@ func (c *SlackChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
opts = append(opts, slack.MsgOptionTS(threadTS))
|
opts = append(opts, slack.MsgOptionTS(threadTS))
|
||||||
}
|
}
|
||||||
|
|
||||||
_, _, err := c.api.PostMessageContext(ctx, channelID, opts...)
|
_, ts, err := c.api.PostMessageContext(ctx, channelID, opts...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("slack send: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("slack send: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
||||||
if ref, ok := c.pendingAcks.LoadAndDelete(msg.ChatID); ok {
|
if ref, ok := c.pendingAcks.LoadAndDelete(msg.ChatID); ok {
|
||||||
|
|
@ -148,23 +148,23 @@ func (c *SlackChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
"thread_ts": threadTS,
|
"thread_ts": threadTS,
|
||||||
})
|
})
|
||||||
|
|
||||||
return nil
|
return []string{ts}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendMedia implements the channels.MediaSender interface.
|
// SendMedia implements the channels.MediaSender interface.
|
||||||
func (c *SlackChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *SlackChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
channelID, _ := parseSlackChatID(msg.ChatID)
|
channelID, _ := parseSlackChatID(msg.ChatID)
|
||||||
if channelID == "" {
|
if channelID == "" {
|
||||||
return fmt.Errorf("invalid slack chat ID: %s", msg.ChatID)
|
return nil, fmt.Errorf("invalid slack chat ID: %s", msg.ChatID)
|
||||||
}
|
}
|
||||||
|
|
||||||
store := c.GetMediaStore()
|
store := c.GetMediaStore()
|
||||||
if store == nil {
|
if store == nil {
|
||||||
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, part := range msg.Parts {
|
for _, part := range msg.Parts {
|
||||||
|
|
@ -198,11 +198,13 @@ func (c *SlackChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessa
|
||||||
"filename": filename,
|
"filename": filename,
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return fmt.Errorf("slack send media: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("slack send media: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
// UploadFileV2 does not expose the posted message timestamp in its
|
||||||
|
// response; returning nil avoids conflating file IDs with message IDs.
|
||||||
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ReactToMessage implements channels.ReactionCapable.
|
// ReactToMessage implements channels.ReactionCapable.
|
||||||
|
|
|
||||||
|
|
@ -16,14 +16,15 @@ func markdownToTelegramHTML(text string) string {
|
||||||
inlineCodes := extractInlineCodes(text)
|
inlineCodes := extractInlineCodes(text)
|
||||||
text = inlineCodes.text
|
text = inlineCodes.text
|
||||||
|
|
||||||
|
links := extractLinks(text)
|
||||||
|
text = links.text
|
||||||
|
|
||||||
text = reHeading.ReplaceAllString(text, "$1")
|
text = reHeading.ReplaceAllString(text, "$1")
|
||||||
|
|
||||||
text = reBlockquote.ReplaceAllString(text, "$1")
|
text = reBlockquote.ReplaceAllString(text, "$1")
|
||||||
|
|
||||||
text = escapeHTML(text)
|
text = escapeHTML(text)
|
||||||
|
|
||||||
text = reLink.ReplaceAllString(text, `<a href="$2">$1</a>`)
|
|
||||||
|
|
||||||
text = reBoldStar.ReplaceAllString(text, "<b>$1</b>")
|
text = reBoldStar.ReplaceAllString(text, "<b>$1</b>")
|
||||||
|
|
||||||
text = reBoldUnder.ReplaceAllString(text, "<b>$1</b>")
|
text = reBoldUnder.ReplaceAllString(text, "<b>$1</b>")
|
||||||
|
|
@ -40,6 +41,12 @@ func markdownToTelegramHTML(text string) string {
|
||||||
|
|
||||||
text = reListItem.ReplaceAllString(text, "• ")
|
text = reListItem.ReplaceAllString(text, "• ")
|
||||||
|
|
||||||
|
for i, lnk := range links.links {
|
||||||
|
label := escapeHTML(lnk[0])
|
||||||
|
url := lnk[1]
|
||||||
|
text = strings.ReplaceAll(text, fmt.Sprintf("\x00LK%d\x00", i), fmt.Sprintf(`<a href="%s">%s</a>`, url, label))
|
||||||
|
}
|
||||||
|
|
||||||
for i, code := range inlineCodes.codes {
|
for i, code := range inlineCodes.codes {
|
||||||
escaped := escapeHTML(code)
|
escaped := escapeHTML(code)
|
||||||
text = strings.ReplaceAll(text, fmt.Sprintf("\x00IC%d\x00", i), fmt.Sprintf("<code>%s</code>", escaped))
|
text = strings.ReplaceAll(text, fmt.Sprintf("\x00IC%d\x00", i), fmt.Sprintf("<code>%s</code>", escaped))
|
||||||
|
|
@ -57,6 +64,29 @@ func markdownToTelegramHTML(text string) string {
|
||||||
return text
|
return text
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type linkMatch struct {
|
||||||
|
text string
|
||||||
|
links [][2]string // [label, url]
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractLinks(text string) linkMatch {
|
||||||
|
matches := reLink.FindAllStringSubmatch(text, -1)
|
||||||
|
|
||||||
|
extracted := make([][2]string, 0, len(matches))
|
||||||
|
for _, match := range matches {
|
||||||
|
extracted = append(extracted, [2]string{match[1], match[2]})
|
||||||
|
}
|
||||||
|
|
||||||
|
i := 0
|
||||||
|
text = reLink.ReplaceAllStringFunc(text, func(m string) string {
|
||||||
|
placeholder := fmt.Sprintf("\x00LK%d\x00", i)
|
||||||
|
i++
|
||||||
|
return placeholder
|
||||||
|
})
|
||||||
|
|
||||||
|
return linkMatch{text: text, links: extracted}
|
||||||
|
}
|
||||||
|
|
||||||
type codeBlockMatch struct {
|
type codeBlockMatch struct {
|
||||||
text string
|
text string
|
||||||
codes []string
|
codes []string
|
||||||
|
|
|
||||||
66
pkg/channels/telegram/parser_markdown_to_html_test.go
Normal file
66
pkg/channels/telegram/parser_markdown_to_html_test.go
Normal file
|
|
@ -0,0 +1,66 @@
|
||||||
|
package telegram
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func Test_markdownToTelegramHTML(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
expected string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "plain text",
|
||||||
|
input: "hello world",
|
||||||
|
expected: "hello world",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "bold",
|
||||||
|
input: "**bold text**",
|
||||||
|
expected: "<b>bold text</b>",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "italic",
|
||||||
|
input: "_italic text_",
|
||||||
|
expected: "<i>italic text</i>",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "link without underscores in URL",
|
||||||
|
input: "[click here](https://example.com/path)",
|
||||||
|
expected: `<a href="https://example.com/path">click here</a>`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "link with underscores in URL is not corrupted by italic regex",
|
||||||
|
// Google Flights URLs use URL-safe base64 with underscores in the tfs param.
|
||||||
|
// Previously reItalic ran after reLink, matching _text_ inside href and injecting
|
||||||
|
// <i> tags into the URL, which broke the link in Telegram.
|
||||||
|
input: "[3 → 10 сентября — от $202](https://www.google.com/travel/flights/search?tfs=CBwQAho_EgoyURL_safe_base64)",
|
||||||
|
expected: `<a href="https://www.google.com/travel/flights/search?tfs=CBwQAho_EgoyURL_safe_base64">3 → 10 сентября — от $202</a>`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "multiple links all survive",
|
||||||
|
input: "[first](https://a.com/path_one) and [second](https://b.com/path_two_x)",
|
||||||
|
expected: `<a href="https://a.com/path_one">first</a> and <a href="https://b.com/path_two_x">second</a>`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "link label with HTML special chars is escaped",
|
||||||
|
input: "[a & b](https://example.com)",
|
||||||
|
expected: `<a href="https://example.com">a & b</a>`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "HTML special chars in plain text are escaped",
|
||||||
|
input: "a & b < c > d",
|
||||||
|
expected: "a & b < c > d",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
actual := markdownToTelegramHTML(tc.input)
|
||||||
|
require.Equal(t, tc.expected, actual)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -168,26 +168,27 @@ func (c *TelegramChannel) Stop(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
useMarkdownV2 := c.config.Channels.Telegram.UseMarkdownV2
|
useMarkdownV2 := c.config.Channels.Telegram.UseMarkdownV2
|
||||||
|
|
||||||
chatID, threadID, err := parseTelegramChatID(msg.ChatID)
|
chatID, threadID, err := parseTelegramChatID(msg.ChatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
return nil, fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
if msg.Content == "" {
|
if msg.Content == "" {
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// The Manager already splits messages to ≤4000 chars (WithMaxMessageLength),
|
// The Manager already splits messages to ≤4000 chars (WithMaxMessageLength),
|
||||||
// so msg.Content is guaranteed to be within that limit. We still need to
|
// so msg.Content is guaranteed to be within that limit. We still need to
|
||||||
// check if HTML expansion pushes it beyond Telegram's 4096-char API limit.
|
// check if HTML expansion pushes it beyond Telegram's 4096-char API limit.
|
||||||
replyToID := msg.ReplyToMessageID
|
replyToID := msg.ReplyToMessageID
|
||||||
|
var messageIDs []string
|
||||||
queue := []string{msg.Content}
|
queue := []string{msg.Content}
|
||||||
for len(queue) > 0 {
|
for len(queue) > 0 {
|
||||||
chunk := queue[0]
|
chunk := queue[0]
|
||||||
|
|
@ -206,16 +207,18 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
}
|
}
|
||||||
|
|
||||||
if smallerLen <= 0 {
|
if smallerLen <= 0 {
|
||||||
if err := c.sendChunk(ctx, sendChunkParams{
|
msgID, err := c.sendChunk(ctx, sendChunkParams{
|
||||||
chatID: chatID,
|
chatID: chatID,
|
||||||
threadID: threadID,
|
threadID: threadID,
|
||||||
content: content,
|
content: content,
|
||||||
replyToID: replyToID,
|
replyToID: replyToID,
|
||||||
mdFallback: chunk,
|
mdFallback: chunk,
|
||||||
useMarkdownV2: useMarkdownV2,
|
useMarkdownV2: useMarkdownV2,
|
||||||
}); err != nil {
|
})
|
||||||
return err
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
|
messageIDs = append(messageIDs, msgID)
|
||||||
replyToID = ""
|
replyToID = ""
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
@ -244,21 +247,23 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.sendChunk(ctx, sendChunkParams{
|
msgID, err := c.sendChunk(ctx, sendChunkParams{
|
||||||
chatID: chatID,
|
chatID: chatID,
|
||||||
threadID: threadID,
|
threadID: threadID,
|
||||||
content: content,
|
content: content,
|
||||||
replyToID: replyToID,
|
replyToID: replyToID,
|
||||||
mdFallback: chunk,
|
mdFallback: chunk,
|
||||||
useMarkdownV2: useMarkdownV2,
|
useMarkdownV2: useMarkdownV2,
|
||||||
}); err != nil {
|
})
|
||||||
return err
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
|
messageIDs = append(messageIDs, msgID)
|
||||||
// Only the first chunk should be a reply; subsequent chunks are normal messages.
|
// Only the first chunk should be a reply; subsequent chunks are normal messages.
|
||||||
replyToID = ""
|
replyToID = ""
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return messageIDs, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type sendChunkParams struct {
|
type sendChunkParams struct {
|
||||||
|
|
@ -275,7 +280,7 @@ type sendChunkParams struct {
|
||||||
func (c *TelegramChannel) sendChunk(
|
func (c *TelegramChannel) sendChunk(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
params sendChunkParams,
|
params sendChunkParams,
|
||||||
) error {
|
) (string, error) {
|
||||||
tgMsg := tu.Message(tu.ID(params.chatID), params.content)
|
tgMsg := tu.Message(tu.ID(params.chatID), params.content)
|
||||||
tgMsg.MessageThreadID = params.threadID
|
tgMsg.MessageThreadID = params.threadID
|
||||||
if params.useMarkdownV2 {
|
if params.useMarkdownV2 {
|
||||||
|
|
@ -292,17 +297,19 @@ func (c *TelegramChannel) sendChunk(
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := c.bot.SendMessage(ctx, tgMsg); err != nil {
|
pMsg, err := c.bot.SendMessage(ctx, tgMsg)
|
||||||
|
if err != nil {
|
||||||
logParseFailed(err, params.useMarkdownV2)
|
logParseFailed(err, params.useMarkdownV2)
|
||||||
|
|
||||||
tgMsg.Text = params.mdFallback
|
tgMsg.Text = params.mdFallback
|
||||||
tgMsg.ParseMode = ""
|
tgMsg.ParseMode = ""
|
||||||
if _, err = c.bot.SendMessage(ctx, tgMsg); err != nil {
|
pMsg, err = c.bot.SendMessage(ctx, tgMsg)
|
||||||
return fmt.Errorf("telegram send: %w", channels.ErrTemporary)
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("telegram send: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return strconv.Itoa(pMsg.MessageID), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// maxTypingDuration limits how long the typing indicator can run.
|
// maxTypingDuration limits how long the typing indicator can run.
|
||||||
|
|
@ -420,21 +427,22 @@ func (c *TelegramChannel) SendPlaceholder(ctx context.Context, chatID string) (s
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendMedia implements the channels.MediaSender interface.
|
// SendMedia implements the channels.MediaSender interface.
|
||||||
func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
chatID, threadID, err := parseTelegramChatID(msg.ChatID)
|
chatID, threadID, err := parseTelegramChatID(msg.ChatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
return nil, fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
store := c.GetMediaStore()
|
store := c.GetMediaStore()
|
||||||
if store == nil {
|
if store == nil {
|
||||||
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var messageIDs []string
|
||||||
for _, part := range msg.Parts {
|
for _, part := range msg.Parts {
|
||||||
localPath, err := store.Resolve(part.Ref)
|
localPath, err := store.Resolve(part.Ref)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -454,6 +462,7 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var tgResult *telego.Message
|
||||||
switch part.Type {
|
switch part.Type {
|
||||||
case "image":
|
case "image":
|
||||||
params := &telego.SendPhotoParams{
|
params := &telego.SendPhotoParams{
|
||||||
|
|
@ -462,11 +471,11 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
Photo: telego.InputFile{File: file},
|
Photo: telego.InputFile{File: file},
|
||||||
Caption: part.Caption,
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
_, err = c.bot.SendPhoto(ctx, params)
|
tgResult, err = c.bot.SendPhoto(ctx, params)
|
||||||
if err != nil && strings.Contains(err.Error(), "PHOTO_INVALID_DIMENSIONS") {
|
if err != nil && strings.Contains(err.Error(), "PHOTO_INVALID_DIMENSIONS") {
|
||||||
if _, seekErr := file.Seek(0, io.SeekStart); seekErr != nil {
|
if _, seekErr := file.Seek(0, io.SeekStart); seekErr != nil {
|
||||||
file.Close()
|
file.Close()
|
||||||
return fmt.Errorf("telegram rewind media after photo failure: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("telegram rewind media after photo failure: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
||||||
docParams := &telego.SendDocumentParams{
|
docParams := &telego.SendDocumentParams{
|
||||||
|
|
@ -475,7 +484,7 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
Document: telego.InputFile{File: file},
|
Document: telego.InputFile{File: file},
|
||||||
Caption: part.Caption,
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
_, err = c.bot.SendDocument(ctx, docParams)
|
tgResult, err = c.bot.SendDocument(ctx, docParams)
|
||||||
}
|
}
|
||||||
case "audio":
|
case "audio":
|
||||||
// Send OGG files with "voice" in the filename as Telegram voice
|
// Send OGG files with "voice" in the filename as Telegram voice
|
||||||
|
|
@ -488,7 +497,7 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
Voice: telego.InputFile{File: file},
|
Voice: telego.InputFile{File: file},
|
||||||
Caption: part.Caption,
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
_, err = c.bot.SendVoice(ctx, vparams)
|
tgResult, err = c.bot.SendVoice(ctx, vparams)
|
||||||
} else {
|
} else {
|
||||||
params := &telego.SendAudioParams{
|
params := &telego.SendAudioParams{
|
||||||
ChatID: tu.ID(chatID),
|
ChatID: tu.ID(chatID),
|
||||||
|
|
@ -496,7 +505,7 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
Audio: telego.InputFile{File: file},
|
Audio: telego.InputFile{File: file},
|
||||||
Caption: part.Caption,
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
_, err = c.bot.SendAudio(ctx, params)
|
tgResult, err = c.bot.SendAudio(ctx, params)
|
||||||
}
|
}
|
||||||
case "video":
|
case "video":
|
||||||
params := &telego.SendVideoParams{
|
params := &telego.SendVideoParams{
|
||||||
|
|
@ -505,7 +514,7 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
Video: telego.InputFile{File: file},
|
Video: telego.InputFile{File: file},
|
||||||
Caption: part.Caption,
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
_, err = c.bot.SendVideo(ctx, params)
|
tgResult, err = c.bot.SendVideo(ctx, params)
|
||||||
default: // "file" or unknown types
|
default: // "file" or unknown types
|
||||||
params := &telego.SendDocumentParams{
|
params := &telego.SendDocumentParams{
|
||||||
ChatID: tu.ID(chatID),
|
ChatID: tu.ID(chatID),
|
||||||
|
|
@ -513,9 +522,12 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
Document: telego.InputFile{File: file},
|
Document: telego.InputFile{File: file},
|
||||||
Caption: part.Caption,
|
Caption: part.Caption,
|
||||||
}
|
}
|
||||||
_, err = c.bot.SendDocument(ctx, params)
|
tgResult, err = c.bot.SendDocument(ctx, params)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if tgResult != nil {
|
||||||
|
messageIDs = append(messageIDs, strconv.Itoa(tgResult.MessageID))
|
||||||
|
}
|
||||||
file.Close()
|
file.Close()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -523,11 +535,11 @@ func (c *TelegramChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMe
|
||||||
"type": part.Type,
|
"type": part.Type,
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return fmt.Errorf("telegram send media: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("telegram send media: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return messageIDs, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Message) error {
|
func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Message) error {
|
||||||
|
|
@ -660,6 +672,23 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
content = cleaned
|
content = cleaned
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if message.ReplyToMessage != nil {
|
||||||
|
quotedMedia := quotedTelegramMediaRefs(
|
||||||
|
message.ReplyToMessage,
|
||||||
|
func(fileID, ext, filename string) string {
|
||||||
|
localPath := c.downloadFile(ctx, fileID, ext)
|
||||||
|
if localPath == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return storeMedia(localPath, filename)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if len(quotedMedia) > 0 {
|
||||||
|
mediaPaths = append(quotedMedia, mediaPaths...)
|
||||||
|
}
|
||||||
|
content = c.prependTelegramQuotedReply(content, message.ReplyToMessage)
|
||||||
|
}
|
||||||
|
|
||||||
// For forum topics, embed the thread ID as "chatID/threadID" so replies
|
// For forum topics, embed the thread ID as "chatID/threadID" so replies
|
||||||
// route to the correct topic and each topic gets its own session.
|
// route to the correct topic and each topic gets its own session.
|
||||||
// Only forum groups (IsForum) are handled; regular group reply threads
|
// Only forum groups (IsForum) are handled; regular group reply threads
|
||||||
|
|
@ -693,6 +722,9 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
"first_name": user.FirstName,
|
"first_name": user.FirstName,
|
||||||
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
||||||
}
|
}
|
||||||
|
if message.ReplyToMessage != nil {
|
||||||
|
metadata["reply_to_message_id"] = fmt.Sprintf("%d", message.ReplyToMessage.MessageID)
|
||||||
|
}
|
||||||
|
|
||||||
// Set parent_peer metadata for per-topic agent binding.
|
// Set parent_peer metadata for per-topic agent binding.
|
||||||
if message.Chat.IsForum && threadID != 0 {
|
if message.Chat.IsForum && threadID != 0 {
|
||||||
|
|
@ -713,6 +745,122 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *TelegramChannel) prependTelegramQuotedReply(content string, reply *telego.Message) string {
|
||||||
|
quoted := strings.TrimSpace(telegramQuotedContent(reply))
|
||||||
|
if quoted == "" {
|
||||||
|
return content
|
||||||
|
}
|
||||||
|
|
||||||
|
author := telegramQuotedAuthor(reply)
|
||||||
|
role := c.telegramQuotedRole(reply)
|
||||||
|
if strings.TrimSpace(content) == "" {
|
||||||
|
return fmt.Sprintf("[quoted %s message from %s]: %s", role, author, quoted)
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("[quoted %s message from %s]: %s\n\n%s", role, author, quoted, content)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TelegramChannel) telegramQuotedRole(message *telego.Message) string {
|
||||||
|
if message == nil {
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
if message.From != nil {
|
||||||
|
if !message.From.IsBot {
|
||||||
|
return "user"
|
||||||
|
}
|
||||||
|
if c.isOwnBotUser(message.From) {
|
||||||
|
return "assistant"
|
||||||
|
}
|
||||||
|
return "bot"
|
||||||
|
}
|
||||||
|
|
||||||
|
if message.SenderChat != nil {
|
||||||
|
return "chat"
|
||||||
|
}
|
||||||
|
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TelegramChannel) isOwnBotUser(user *telego.User) bool {
|
||||||
|
if c == nil || c.bot == nil || user == nil || !user.IsBot {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if botID := c.bot.ID(); botID != 0 && user.ID == botID {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
botUsername := strings.TrimPrefix(strings.TrimSpace(c.bot.Username()), "@")
|
||||||
|
if botUsername == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return strings.EqualFold(strings.TrimPrefix(strings.TrimSpace(user.Username), "@"), botUsername)
|
||||||
|
}
|
||||||
|
|
||||||
|
func telegramQuotedAuthor(message *telego.Message) string {
|
||||||
|
if message == nil || message.From == nil {
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
if username := strings.TrimSpace(message.From.Username); username != "" {
|
||||||
|
return username
|
||||||
|
}
|
||||||
|
if firstName := strings.TrimSpace(message.From.FirstName); firstName != "" {
|
||||||
|
return firstName
|
||||||
|
}
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
func telegramQuotedContent(message *telego.Message) string {
|
||||||
|
if message == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var parts []string
|
||||||
|
if text := strings.TrimSpace(message.Text); text != "" {
|
||||||
|
parts = append(parts, text)
|
||||||
|
}
|
||||||
|
if caption := strings.TrimSpace(message.Caption); caption != "" {
|
||||||
|
parts = append(parts, caption)
|
||||||
|
}
|
||||||
|
switch {
|
||||||
|
case len(message.Photo) > 0:
|
||||||
|
parts = append(parts, "[image: photo]")
|
||||||
|
}
|
||||||
|
switch {
|
||||||
|
case message.Voice != nil:
|
||||||
|
parts = append(parts, "[voice]")
|
||||||
|
case message.Audio != nil:
|
||||||
|
parts = append(parts, "[audio]")
|
||||||
|
}
|
||||||
|
if message.Document != nil {
|
||||||
|
parts = append(parts, "[file]")
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.Join(parts, "\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
func quotedTelegramMediaRefs(
|
||||||
|
message *telego.Message,
|
||||||
|
resolve func(fileID, ext, filename string) string,
|
||||||
|
) []string {
|
||||||
|
if message == nil || resolve == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var refs []string
|
||||||
|
if message.Voice != nil {
|
||||||
|
if ref := resolve(message.Voice.FileID, ".ogg", "voice.ogg"); ref != "" {
|
||||||
|
refs = append(refs, ref)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if message.Audio != nil {
|
||||||
|
if ref := resolve(message.Audio.FileID, ".mp3", "audio.mp3"); ref != "" {
|
||||||
|
refs = append(refs, ref)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return refs
|
||||||
|
}
|
||||||
|
|
||||||
func (c *TelegramChannel) downloadPhoto(ctx context.Context, fileID string) string {
|
func (c *TelegramChannel) downloadPhoto(ctx context.Context, fileID string) string {
|
||||||
file, err := c.bot.GetFile(ctx, &telego.GetFileParams{FileID: fileID})
|
file, err := c.bot.GetFile(ctx, &telego.GetFileParams{FileID: fileID})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
|
||||||
|
|
@ -7,6 +7,7 @@ import (
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
|
@ -104,6 +105,13 @@ func successResponse(t *testing.T) *ta.Response {
|
||||||
return &ta.Response{Ok: true, Result: b}
|
return &ta.Response{Ok: true, Result: b}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func successUserResponse(t *testing.T, user *telego.User) *ta.Response {
|
||||||
|
t.Helper()
|
||||||
|
b, err := json.Marshal(user)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return &ta.Response{Ok: true, Result: b}
|
||||||
|
}
|
||||||
|
|
||||||
// newTestChannel creates a TelegramChannel with a mocked bot for unit testing.
|
// newTestChannel creates a TelegramChannel with a mocked bot for unit testing.
|
||||||
func newTestChannel(t *testing.T, caller *stubCaller) *TelegramChannel {
|
func newTestChannel(t *testing.T, caller *stubCaller) *TelegramChannel {
|
||||||
return newTestChannelWithConstructor(t, caller, &stubConstructor{})
|
return newTestChannelWithConstructor(t, caller, &stubConstructor{})
|
||||||
|
|
@ -168,7 +176,7 @@ func TestSendMedia_ImageFallbacksToDocumentOnInvalidDimensions(t *testing.T) {
|
||||||
)
|
)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
ChatID: "12345",
|
ChatID: "12345",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
Type: "image",
|
Type: "image",
|
||||||
|
|
@ -206,7 +214,7 @@ func TestSendMedia_ImageNonDimensionErrorDoesNotFallback(t *testing.T) {
|
||||||
ref, err := store.Store(localPath, media.MediaMeta{Filename: "image.png", ContentType: "image/png"}, "scope-1")
|
ref, err := store.Store(localPath, media.MediaMeta{Filename: "image.png", ContentType: "image/png"}, "scope-1")
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
ChatID: "12345",
|
ChatID: "12345",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
Type: "image",
|
Type: "image",
|
||||||
|
|
@ -231,7 +239,7 @@ func TestSend_EmptyContent(t *testing.T) {
|
||||||
}
|
}
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
ChatID: "12345",
|
ChatID: "12345",
|
||||||
Content: "",
|
Content: "",
|
||||||
})
|
})
|
||||||
|
|
@ -248,7 +256,7 @@ func TestSend_ShortMessage_SingleCall(t *testing.T) {
|
||||||
}
|
}
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
ChatID: "12345",
|
ChatID: "12345",
|
||||||
Content: "Hello, world!",
|
Content: "Hello, world!",
|
||||||
})
|
})
|
||||||
|
|
@ -271,7 +279,7 @@ func TestSend_LongMessage_SingleCall(t *testing.T) {
|
||||||
|
|
||||||
longContent := strings.Repeat("a", 4000)
|
longContent := strings.Repeat("a", 4000)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
ChatID: "12345",
|
ChatID: "12345",
|
||||||
Content: longContent,
|
Content: longContent,
|
||||||
})
|
})
|
||||||
|
|
@ -294,7 +302,7 @@ func TestSend_HTMLFallback_PerChunk(t *testing.T) {
|
||||||
}
|
}
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
ChatID: "12345",
|
ChatID: "12345",
|
||||||
Content: "Hello **world**",
|
Content: "Hello **world**",
|
||||||
})
|
})
|
||||||
|
|
@ -312,7 +320,7 @@ func TestSend_HTMLFallback_BothFail(t *testing.T) {
|
||||||
}
|
}
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
ChatID: "12345",
|
ChatID: "12345",
|
||||||
Content: "Hello",
|
Content: "Hello",
|
||||||
})
|
})
|
||||||
|
|
@ -334,7 +342,7 @@ func TestSend_LongMessage_HTMLFallback_StopsOnError(t *testing.T) {
|
||||||
|
|
||||||
longContent := strings.Repeat("x", 4001)
|
longContent := strings.Repeat("x", 4001)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
ChatID: "12345",
|
ChatID: "12345",
|
||||||
Content: longContent,
|
Content: longContent,
|
||||||
})
|
})
|
||||||
|
|
@ -364,7 +372,7 @@ func TestSend_MarkdownShortButHTMLLong_MultipleCalls(t *testing.T) {
|
||||||
"HTML expansion must exceed Telegram limit for this test to be meaningful",
|
"HTML expansion must exceed Telegram limit for this test to be meaningful",
|
||||||
)
|
)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
ChatID: "12345",
|
ChatID: "12345",
|
||||||
Content: markdownContent,
|
Content: markdownContent,
|
||||||
})
|
})
|
||||||
|
|
@ -399,7 +407,7 @@ func TestSend_HTMLOverflow_WordBoundary(t *testing.T) {
|
||||||
// Ensure the test content matches the intended boundary conditions.
|
// Ensure the test content matches the intended boundary conditions.
|
||||||
assert.LessOrEqual(t, len([]rune(content)), 4000, "markdown content must not exceed chunk size for this test")
|
assert.LessOrEqual(t, len([]rune(content)), 4000, "markdown content must not exceed chunk size for this test")
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
ChatID: "123456",
|
ChatID: "123456",
|
||||||
Content: content,
|
Content: content,
|
||||||
})
|
})
|
||||||
|
|
@ -435,7 +443,7 @@ func TestSend_NotRunning(t *testing.T) {
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
ch.SetRunning(false)
|
ch.SetRunning(false)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
ChatID: "12345",
|
ChatID: "12345",
|
||||||
Content: "Hello",
|
Content: "Hello",
|
||||||
})
|
})
|
||||||
|
|
@ -453,7 +461,7 @@ func TestSend_InvalidChatID(t *testing.T) {
|
||||||
}
|
}
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
ChatID: "not-a-number",
|
ChatID: "not-a-number",
|
||||||
Content: "Hello",
|
Content: "Hello",
|
||||||
})
|
})
|
||||||
|
|
@ -510,7 +518,7 @@ func TestSend_WithForumThreadID(t *testing.T) {
|
||||||
}
|
}
|
||||||
ch := newTestChannel(t, caller)
|
ch := newTestChannel(t, caller)
|
||||||
|
|
||||||
err := ch.Send(context.Background(), bus.OutboundMessage{
|
_, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
ChatID: "-1001234567890/42",
|
ChatID: "-1001234567890/42",
|
||||||
Content: "Hello from topic",
|
Content: "Hello from topic",
|
||||||
})
|
})
|
||||||
|
|
@ -642,6 +650,181 @@ func TestHandleMessage_ReplyThread_NonForum_NoIsolation(t *testing.T) {
|
||||||
assert.Empty(t, inbound.Metadata["parent_peer_id"])
|
assert.Empty(t, inbound.Metadata["parent_peer_id"])
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func assertHandleMessageQuotedUserReply(
|
||||||
|
t *testing.T,
|
||||||
|
chatID int64,
|
||||||
|
messageID int,
|
||||||
|
userID int64,
|
||||||
|
userName string,
|
||||||
|
userText string,
|
||||||
|
replyMessageID int,
|
||||||
|
replyText string,
|
||||||
|
replyCaption string,
|
||||||
|
replyAuthorID int64,
|
||||||
|
replyAuthorName string,
|
||||||
|
expectedContent string,
|
||||||
|
) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
messageBus := bus.NewMessageBus()
|
||||||
|
ch := &TelegramChannel{
|
||||||
|
BaseChannel: channels.NewBaseChannel("telegram", nil, messageBus, nil),
|
||||||
|
chatIDs: make(map[string]int64),
|
||||||
|
ctx: context.Background(),
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := &telego.Message{
|
||||||
|
Text: userText,
|
||||||
|
MessageID: messageID,
|
||||||
|
Chat: telego.Chat{
|
||||||
|
ID: chatID,
|
||||||
|
Type: "private",
|
||||||
|
},
|
||||||
|
From: &telego.User{
|
||||||
|
ID: userID,
|
||||||
|
FirstName: userName,
|
||||||
|
},
|
||||||
|
ReplyToMessage: &telego.Message{
|
||||||
|
MessageID: replyMessageID,
|
||||||
|
Text: replyText,
|
||||||
|
Caption: replyCaption,
|
||||||
|
From: &telego.User{
|
||||||
|
ID: replyAuthorID,
|
||||||
|
FirstName: replyAuthorName,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
err := ch.handleMessage(context.Background(), msg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
inbound, ok := <-messageBus.InboundChan()
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.Equal(t, strconv.Itoa(replyMessageID), inbound.Metadata["reply_to_message_id"])
|
||||||
|
assert.Equal(t, expectedContent, inbound.Content)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleMessage_ReplyToMessage_PrependsQuotedTextAndMetadata(t *testing.T) {
|
||||||
|
assertHandleMessageQuotedUserReply(
|
||||||
|
t,
|
||||||
|
456,
|
||||||
|
21,
|
||||||
|
11,
|
||||||
|
"Alice",
|
||||||
|
"follow up",
|
||||||
|
99,
|
||||||
|
"old context",
|
||||||
|
"",
|
||||||
|
12,
|
||||||
|
"Bob",
|
||||||
|
"[quoted user message from Bob]: old context\n\nfollow up",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleMessage_ReplyToMessage_UsesCaptionWhenQuotedTextMissing(t *testing.T) {
|
||||||
|
assertHandleMessageQuotedUserReply(
|
||||||
|
t,
|
||||||
|
789,
|
||||||
|
22,
|
||||||
|
13,
|
||||||
|
"Carol",
|
||||||
|
"answer this",
|
||||||
|
100,
|
||||||
|
"",
|
||||||
|
"caption context",
|
||||||
|
14,
|
||||||
|
"Dave",
|
||||||
|
"[quoted user message from Dave]: caption context\n\nanswer this",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleMessage_ReplyToOwnBotMessage_UsesAssistantRole(t *testing.T) {
|
||||||
|
messageBus := bus.NewMessageBus()
|
||||||
|
caller := &stubCaller{
|
||||||
|
callFn: func(ctx context.Context, url string, data *ta.RequestData) (*ta.Response, error) {
|
||||||
|
if strings.Contains(url, "getMe") {
|
||||||
|
return successUserResponse(t, &telego.User{
|
||||||
|
ID: 42,
|
||||||
|
IsBot: true,
|
||||||
|
FirstName: "Pico",
|
||||||
|
Username: "afjcjsbx_picoclaw_bot",
|
||||||
|
}), nil
|
||||||
|
}
|
||||||
|
t.Fatalf("unexpected API call: %s", url)
|
||||||
|
return nil, nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
ch := newTestChannel(t, caller)
|
||||||
|
ch.BaseChannel = channels.NewBaseChannel("telegram", nil, messageBus, nil)
|
||||||
|
ch.ctx = context.Background()
|
||||||
|
|
||||||
|
msg := &telego.Message{
|
||||||
|
Text: "ti ricordi questo file?",
|
||||||
|
MessageID: 23,
|
||||||
|
Chat: telego.Chat{
|
||||||
|
ID: 999,
|
||||||
|
Type: "private",
|
||||||
|
},
|
||||||
|
From: &telego.User{
|
||||||
|
ID: 15,
|
||||||
|
FirstName: "Eve",
|
||||||
|
},
|
||||||
|
ReplyToMessage: &telego.Message{
|
||||||
|
MessageID: 101,
|
||||||
|
Text: "Fatto! Ho creato il file notizie_2026_03_28.md",
|
||||||
|
From: &telego.User{
|
||||||
|
ID: 42,
|
||||||
|
IsBot: true,
|
||||||
|
FirstName: "Pico",
|
||||||
|
Username: "afjcjsbx_picoclaw_bot",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
err := ch.handleMessage(context.Background(), msg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
inbound, ok := <-messageBus.InboundChan()
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.Equal(t, "101", inbound.Metadata["reply_to_message_id"])
|
||||||
|
assert.Equal(
|
||||||
|
t,
|
||||||
|
"[quoted assistant message from afjcjsbx_picoclaw_bot]: Fatto! Ho creato il file notizie_2026_03_28.md\n\nti ricordi questo file?",
|
||||||
|
inbound.Content,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTelegramQuotedContent_IncludesVoiceMarkerAlongsideCaption(t *testing.T) {
|
||||||
|
msg := &telego.Message{
|
||||||
|
Caption: "listen to this",
|
||||||
|
Voice: &telego.Voice{
|
||||||
|
FileID: "voice-file",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.Equal(t, "listen to this\n[voice]", telegramQuotedContent(msg))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQuotedTelegramMediaRefs_ResolvesQuotedAudioInOrder(t *testing.T) {
|
||||||
|
msg := &telego.Message{
|
||||||
|
Voice: &telego.Voice{FileID: "voice-file"},
|
||||||
|
Audio: &telego.Audio{FileID: "audio-file"},
|
||||||
|
}
|
||||||
|
|
||||||
|
var calls []string
|
||||||
|
refs := quotedTelegramMediaRefs(msg, func(fileID, ext, filename string) string {
|
||||||
|
calls = append(calls, fileID+"|"+ext+"|"+filename)
|
||||||
|
return "ref://" + filename
|
||||||
|
})
|
||||||
|
|
||||||
|
assert.Equal(
|
||||||
|
t,
|
||||||
|
[]string{"voice-file|.ogg|voice.ogg", "audio-file|.mp3|audio.mp3"},
|
||||||
|
calls,
|
||||||
|
)
|
||||||
|
assert.Equal(t, []string{"ref://voice.ogg", "ref://audio.mp3"}, refs)
|
||||||
|
}
|
||||||
|
|
||||||
func TestHandleMessage_EmptyContent_Ignored(t *testing.T) {
|
func TestHandleMessage_EmptyContent_Ignored(t *testing.T) {
|
||||||
messageBus := bus.NewMessageBus()
|
messageBus := bus.NewMessageBus()
|
||||||
ch := &TelegramChannel{
|
ch := &TelegramChannel{
|
||||||
|
|
|
||||||
|
|
@ -184,20 +184,20 @@ func (c *WeComChannel) BeginStream(_ context.Context, chatID string) (channels.S
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *WeComChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *WeComChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
content := strings.TrimSpace(msg.Content)
|
content := strings.TrimSpace(msg.Content)
|
||||||
if content == "" {
|
if content == "" {
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if turn, ok := c.getTurn(msg.ChatID); ok {
|
if turn, ok := c.getTurn(msg.ChatID); ok {
|
||||||
if time.Since(turn.CreatedAt) <= wecomStreamMaxDuration {
|
if time.Since(turn.CreatedAt) <= wecomStreamMaxDuration {
|
||||||
if err := c.sendStreamReply(turn, content); err == nil {
|
if err := c.sendStreamReply(turn, content); err == nil {
|
||||||
c.consumeTurn(msg.ChatID, turn)
|
c.consumeTurn(msg.ChatID, turn)
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
c.consumeTurn(msg.ChatID, turn)
|
c.consumeTurn(msg.ChatID, turn)
|
||||||
|
|
@ -205,20 +205,20 @@ func (c *WeComChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
|
|
||||||
if route, ok := c.routes.Get(msg.ChatID); ok {
|
if route, ok := c.routes.Get(msg.ChatID); ok {
|
||||||
if err := c.sendActivePush(route.ChatID, route.ChatType, content); err != nil {
|
if err := c.sendActivePush(route.ChatID, route.ChatType, content); err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.sendActivePush(msg.ChatID, 0, content); err != nil {
|
if err := c.sendActivePush(msg.ChatID, 0, content); err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *WeComChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *WeComChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
route, chatType, hasTurn := c.resolveMediaRoute(msg.ChatID)
|
route, chatType, hasTurn := c.resolveMediaRoute(msg.ChatID)
|
||||||
|
|
@ -231,7 +231,7 @@ func (c *WeComChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessa
|
||||||
if strings.TrimSpace(part.Ref) == "" {
|
if strings.TrimSpace(part.Ref) == "" {
|
||||||
if caption := strings.TrimSpace(part.Caption); caption != "" {
|
if caption := strings.TrimSpace(part.Caption); caption != "" {
|
||||||
if err := c.sendActivePush(chatID, chatType, caption); err != nil {
|
if err := c.sendActivePush(chatID, chatType, caption); err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
|
|
@ -239,7 +239,7 @@ func (c *WeComChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessa
|
||||||
|
|
||||||
localPath, filename, contentType, cleanup, err := c.resolveOutboundPart(ctx, part)
|
localPath, filename, contentType, cleanup, err := c.resolveOutboundPart(ctx, part)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("wecom resolve media %q: %v: %w", part.Ref, err, channels.ErrSendFailed)
|
return nil, fmt.Errorf("wecom resolve media %q: %v: %w", part.Ref, err, channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
||||||
func() {
|
func() {
|
||||||
|
|
@ -283,11 +283,11 @@ func (c *WeComChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessa
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *WeComChannel) connectLoop() {
|
func (c *WeComChannel) connectLoop() {
|
||||||
|
|
|
||||||
|
|
@ -190,7 +190,7 @@ func TestSend_StreamFailureFallsBackToActualChatID(t *testing.T) {
|
||||||
return wecomTestAck(nil), nil
|
return wecomTestAck(nil), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := ch.Send(context.Background(), bus.OutboundMessage{
|
if _, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
Channel: "wecom",
|
Channel: "wecom",
|
||||||
ChatID: "chat-1",
|
ChatID: "chat-1",
|
||||||
Content: "hello",
|
Content: "hello",
|
||||||
|
|
@ -247,7 +247,7 @@ func TestSend_DoesNotSplitStreamReply(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
content := strings.Repeat("\u4e2d", 30000)
|
content := strings.Repeat("\u4e2d", 30000)
|
||||||
if err := ch.Send(context.Background(), bus.OutboundMessage{
|
if _, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
Channel: "wecom",
|
Channel: "wecom",
|
||||||
ChatID: "chat-1",
|
ChatID: "chat-1",
|
||||||
Content: content,
|
Content: content,
|
||||||
|
|
@ -283,7 +283,7 @@ func TestSend_DoesNotSplitActivePush(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
content := strings.Repeat("a", 30000)
|
content := strings.Repeat("a", 30000)
|
||||||
if err := ch.Send(context.Background(), bus.OutboundMessage{
|
if _, err := ch.Send(context.Background(), bus.OutboundMessage{
|
||||||
Channel: "wecom",
|
Channel: "wecom",
|
||||||
ChatID: "chat-1",
|
ChatID: "chat-1",
|
||||||
Content: content,
|
Content: content,
|
||||||
|
|
@ -346,7 +346,7 @@ func TestSendMedia_SendsActiveImage(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
Channel: "wecom",
|
Channel: "wecom",
|
||||||
ChatID: "chat-1",
|
ChatID: "chat-1",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
|
|
@ -457,7 +457,7 @@ func TestSendMedia_UsesTurnImageAndFinishesStream(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
Channel: "wecom",
|
Channel: "wecom",
|
||||||
ChatID: "chat-1",
|
ChatID: "chat-1",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
|
|
@ -553,7 +553,7 @@ func TestSendMedia_SendsActiveFile(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
|
||||||
Channel: "wecom",
|
Channel: "wecom",
|
||||||
ChatID: "chat-2",
|
ChatID: "chat-2",
|
||||||
Parts: []bus.MediaPart{{
|
Parts: []bus.MediaPart{{
|
||||||
|
|
|
||||||
|
|
@ -1097,12 +1097,12 @@ func (c *WeixinChannel) StartTyping(ctx context.Context, chatID string) (func(),
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendMedia implements channels.MediaSender.
|
// SendMedia implements channels.MediaSender.
|
||||||
func (c *WeixinChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
func (c *WeixinChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return basechannels.ErrNotRunning
|
return nil, basechannels.ErrNotRunning
|
||||||
}
|
}
|
||||||
if err := c.ensureSessionActive(); err != nil {
|
if err := c.ensureSessionActive(); err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
contextToken := ""
|
contextToken := ""
|
||||||
|
|
@ -1110,7 +1110,7 @@ func (c *WeixinChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
contextToken, _ = v.(string)
|
contextToken, _ = v.(string)
|
||||||
}
|
}
|
||||||
if contextToken == "" {
|
if contextToken == "" {
|
||||||
return fmt.Errorf(
|
return nil, fmt.Errorf(
|
||||||
"weixin send media: missing context token for chat %s: %w",
|
"weixin send media: missing context token for chat %s: %w",
|
||||||
msg.ChatID,
|
msg.ChatID,
|
||||||
basechannels.ErrSendFailed,
|
basechannels.ErrSendFailed,
|
||||||
|
|
@ -1125,7 +1125,7 @@ func (c *WeixinChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
"ref": part.Ref,
|
"ref": part.Ref,
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return fmt.Errorf("weixin send media: %w", basechannels.ErrSendFailed)
|
return nil, fmt.Errorf("weixin send media: %w", basechannels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
func() {
|
func() {
|
||||||
if cleanup != nil {
|
if cleanup != nil {
|
||||||
|
|
@ -1147,11 +1147,11 @@ func (c *WeixinChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMess
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
if c.remainingPause() > 0 {
|
if c.remainingPause() > 0 {
|
||||||
return fmt.Errorf("weixin send media: %w", basechannels.ErrSendFailed)
|
return nil, fmt.Errorf("weixin send media: %w", basechannels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
return fmt.Errorf("weixin send media: %w", basechannels.ErrTemporary)
|
return nil, fmt.Errorf("weixin send media: %w", basechannels.ErrTemporary)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -358,16 +358,16 @@ func (c *WeixinChannel) handleInboundMessage(ctx context.Context, msg WeixinMess
|
||||||
}
|
}
|
||||||
|
|
||||||
// Send implements channels.Channel by sending a text message to the WeChat user.
|
// Send implements channels.Channel by sending a text message to the WeChat user.
|
||||||
func (c *WeixinChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *WeixinChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
if err := c.ensureSessionActive(); err != nil {
|
if err := c.ensureSessionActive(); err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
if msg.Content == "" {
|
if msg.Content == "" {
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// We need a context_token to send a reply. It should be stored in the conversation metadata.
|
// We need a context_token to send a reply. It should be stored in the conversation metadata.
|
||||||
|
|
@ -386,7 +386,7 @@ func (c *WeixinChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
logger.ErrorCF("weixin", "Missing context token, cannot send message", map[string]any{
|
logger.ErrorCF("weixin", "Missing context token, cannot send message", map[string]any{
|
||||||
"to_user_id": toUserID,
|
"to_user_id": toUserID,
|
||||||
})
|
})
|
||||||
return fmt.Errorf("weixin send: %w: missing context token for chat %s", channels.ErrSendFailed, toUserID)
|
return nil, fmt.Errorf("weixin send: %w: missing context token for chat %s", channels.ErrSendFailed, toUserID)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.sendTextMessage(ctx, toUserID, contextToken, msg.Content); err != nil {
|
if err := c.sendTextMessage(ctx, toUserID, contextToken, msg.Content); err != nil {
|
||||||
|
|
@ -395,10 +395,10 @@ func (c *WeixinChannel) Send(ctx context.Context, msg bus.OutboundMessage) error
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
if c.remainingPause() > 0 {
|
if c.remainingPause() > 0 {
|
||||||
return fmt.Errorf("weixin send: %w", channels.ErrSendFailed)
|
return nil, fmt.Errorf("weixin send: %w", channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
return fmt.Errorf("weixin send: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("weixin send: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -104,15 +104,15 @@ func (c *WhatsAppChannel) Stop(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *WhatsAppChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *WhatsAppChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check ctx before acquiring lock
|
// Check ctx before acquiring lock
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return ctx.Err()
|
return nil, ctx.Err()
|
||||||
default:
|
default:
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -120,7 +120,7 @@ func (c *WhatsAppChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
defer c.mu.Unlock()
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
if c.conn == nil {
|
if c.conn == nil {
|
||||||
return fmt.Errorf("whatsapp connection not established: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("whatsapp connection not established: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
||||||
payload := map[string]any{
|
payload := map[string]any{
|
||||||
|
|
@ -131,17 +131,17 @@ func (c *WhatsAppChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
|
|
||||||
data, err := json.Marshal(payload)
|
data, err := json.Marshal(payload)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to marshal message: %w", err)
|
return nil, fmt.Errorf("failed to marshal message: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
_ = c.conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
_ = c.conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||||
if err := c.conn.WriteMessage(websocket.TextMessage, data); err != nil {
|
if err := c.conn.WriteMessage(websocket.TextMessage, data); err != nil {
|
||||||
_ = c.conn.SetWriteDeadline(time.Time{})
|
_ = c.conn.SetWriteDeadline(time.Time{})
|
||||||
return fmt.Errorf("whatsapp send: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("whatsapp send: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
_ = c.conn.SetWriteDeadline(time.Time{})
|
_ = c.conn.SetWriteDeadline(time.Time{})
|
||||||
|
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *WhatsAppChannel) listen() {
|
func (c *WhatsAppChannel) listen() {
|
||||||
|
|
|
||||||
|
|
@ -396,13 +396,13 @@ func (c *WhatsAppNativeChannel) handleIncoming(evt *events.Message) {
|
||||||
c.HandleMessage(c.runCtx, peer, messageID, senderID, chatID, content, mediaPaths, metadata, sender)
|
c.HandleMessage(c.runCtx, peer, messageID, senderID, chatID, content, mediaPaths, metadata, sender)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *WhatsAppNativeChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *WhatsAppNativeChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return nil, channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return ctx.Err()
|
return nil, ctx.Err()
|
||||||
default:
|
default:
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -411,18 +411,18 @@ func (c *WhatsAppNativeChannel) Send(ctx context.Context, msg bus.OutboundMessag
|
||||||
c.mu.Unlock()
|
c.mu.Unlock()
|
||||||
|
|
||||||
if client == nil || !client.IsConnected() {
|
if client == nil || !client.IsConnected() {
|
||||||
return fmt.Errorf("whatsapp connection not established: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("whatsapp connection not established: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Detect unpaired state: the client is connected (to WhatsApp servers)
|
// Detect unpaired state: the client is connected (to WhatsApp servers)
|
||||||
// but has not completed QR-login yet, so sending would fail.
|
// but has not completed QR-login yet, so sending would fail.
|
||||||
if client.Store.ID == nil {
|
if client.Store.ID == nil {
|
||||||
return fmt.Errorf("whatsapp not yet paired (QR login pending): %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("whatsapp not yet paired (QR login pending): %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
|
|
||||||
to, err := parseJID(msg.ChatID)
|
to, err := parseJID(msg.ChatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid chat id %q: %w", msg.ChatID, err)
|
return nil, fmt.Errorf("invalid chat id %q: %w", msg.ChatID, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
waMsg := &waE2E.Message{
|
waMsg := &waE2E.Message{
|
||||||
|
|
@ -430,9 +430,9 @@ func (c *WhatsAppNativeChannel) Send(ctx context.Context, msg bus.OutboundMessag
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err = client.SendMessage(ctx, to, waMsg); err != nil {
|
if _, err = client.SendMessage(ctx, to, waMsg); err != nil {
|
||||||
return fmt.Errorf("whatsapp send: %w", channels.ErrTemporary)
|
return nil, fmt.Errorf("whatsapp send: %w", channels.ErrTemporary)
|
||||||
}
|
}
|
||||||
return nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// parseJID converts a chat ID (phone number or JID string) to types.JID.
|
// parseJID converts a chat ID (phone number or JID string) to types.JID.
|
||||||
|
|
|
||||||
|
|
@ -208,7 +208,10 @@ func (p *Provider) Chat(
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Check for SSO token expiration errors and provide actionable guidance
|
// Check for SSO token expiration errors and provide actionable guidance
|
||||||
if isSSOTokenError(err) {
|
if isSSOTokenError(err) {
|
||||||
return nil, fmt.Errorf("bedrock converse: AWS credentials may have expired. If using AWS SSO, run 'aws sso login' to refresh: %w", err)
|
return nil, fmt.Errorf(
|
||||||
|
"bedrock converse: AWS credentials may have expired. If using AWS SSO, run 'aws sso login' to refresh: %w",
|
||||||
|
err,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
return nil, fmt.Errorf("bedrock converse: %w", err)
|
return nil, fmt.Errorf("bedrock converse: %w", err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -583,13 +583,17 @@ func TestIsSSOTokenError(t *testing.T) {
|
||||||
expected: true,
|
expected: true,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "full SSO error message",
|
name: "full SSO error message",
|
||||||
err: fmt.Errorf("get identity: get credentials: failed to refresh cached credentials, refresh cached SSO token failed, unable to refresh SSO token"),
|
err: fmt.Errorf(
|
||||||
|
"get identity: get credentials: failed to refresh cached credentials, refresh cached SSO token failed, unable to refresh SSO token",
|
||||||
|
),
|
||||||
expected: true,
|
expected: true,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "SSO token file missing",
|
name: "SSO token file missing",
|
||||||
err: fmt.Errorf("get identity: get credentials: failed to refresh cached credentials, failed to read cached SSO token file, open ~/.aws/sso/cache/abc123.json: no such file or directory"),
|
err: fmt.Errorf(
|
||||||
|
"get identity: get credentials: failed to refresh cached credentials, failed to read cached SSO token file, open ~/.aws/sso/cache/abc123.json: no such file or directory",
|
||||||
|
),
|
||||||
expected: true,
|
expected: true,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -17,6 +17,48 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/providers/bedrock"
|
"github.com/sipeed/picoclaw/pkg/providers/bedrock"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type protocolMeta struct {
|
||||||
|
defaultAPIBase string
|
||||||
|
emptyAPIKeyAllowed bool
|
||||||
|
}
|
||||||
|
|
||||||
|
var protocolMetaByName = map[string]protocolMeta{
|
||||||
|
"openai": {defaultAPIBase: "https://api.openai.com/v1"},
|
||||||
|
"openrouter": {defaultAPIBase: "https://openrouter.ai/api/v1"},
|
||||||
|
"litellm": {defaultAPIBase: "http://localhost:4000/v1"},
|
||||||
|
"lmstudio": {defaultAPIBase: "http://localhost:1234/v1", emptyAPIKeyAllowed: true},
|
||||||
|
"novita": {defaultAPIBase: "https://api.novita.ai/openai"},
|
||||||
|
"groq": {defaultAPIBase: "https://api.groq.com/openai/v1"},
|
||||||
|
"zhipu": {defaultAPIBase: "https://open.bigmodel.cn/api/paas/v4"},
|
||||||
|
"gemini": {defaultAPIBase: "https://generativelanguage.googleapis.com/v1beta"},
|
||||||
|
"nvidia": {defaultAPIBase: "https://integrate.api.nvidia.com/v1"},
|
||||||
|
"ollama": {defaultAPIBase: "http://localhost:11434/v1", emptyAPIKeyAllowed: true},
|
||||||
|
"moonshot": {defaultAPIBase: "https://api.moonshot.cn/v1"},
|
||||||
|
"shengsuanyun": {defaultAPIBase: "https://router.shengsuanyun.com/api/v1"},
|
||||||
|
"deepseek": {defaultAPIBase: "https://api.deepseek.com/v1"},
|
||||||
|
"cerebras": {defaultAPIBase: "https://api.cerebras.ai/v1"},
|
||||||
|
"vivgrid": {defaultAPIBase: "https://api.vivgrid.com/v1"},
|
||||||
|
"volcengine": {defaultAPIBase: "https://ark.cn-beijing.volces.com/api/v3"},
|
||||||
|
"qwen": {defaultAPIBase: "https://dashscope.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"qwen-intl": {defaultAPIBase: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"qwen-international": {defaultAPIBase: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"dashscope-intl": {defaultAPIBase: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"qwen-us": {defaultAPIBase: "https://dashscope-us.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"dashscope-us": {defaultAPIBase: "https://dashscope-us.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"coding-plan": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/v1"},
|
||||||
|
"alibaba-coding": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/v1"},
|
||||||
|
"qwen-coding": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/v1"},
|
||||||
|
"coding-plan-anthropic": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/apps/anthropic"},
|
||||||
|
"alibaba-coding-anthropic": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/apps/anthropic"},
|
||||||
|
"vllm": {defaultAPIBase: "http://localhost:8000/v1", emptyAPIKeyAllowed: true},
|
||||||
|
"mistral": {defaultAPIBase: "https://api.mistral.ai/v1"},
|
||||||
|
"avian": {defaultAPIBase: "https://api.avian.io/v1"},
|
||||||
|
"minimax": {defaultAPIBase: "https://api.minimaxi.com/v1"},
|
||||||
|
"longcat": {defaultAPIBase: "https://api.longcat.chat/openai"},
|
||||||
|
"modelscope": {defaultAPIBase: "https://api-inference.modelscope.cn/v1"},
|
||||||
|
"mimo": {defaultAPIBase: "https://api.xiaomimimo.com/v1"},
|
||||||
|
}
|
||||||
|
|
||||||
// createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store.
|
// createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store.
|
||||||
func createClaudeAuthProvider() (LLMProvider, error) {
|
func createClaudeAuthProvider() (LLMProvider, error) {
|
||||||
cred, err := getCredential("anthropic")
|
cred, err := getCredential("anthropic")
|
||||||
|
|
@ -154,13 +196,13 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
|
||||||
}
|
}
|
||||||
return provider, modelID, nil
|
return provider, modelID, nil
|
||||||
|
|
||||||
case "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
case "litellm", "lmstudio", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
||||||
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||||
"vivgrid", "volcengine", "vllm", "qwen", "qwen-intl", "qwen-international", "dashscope-intl",
|
"vivgrid", "volcengine", "vllm", "qwen", "qwen-intl", "qwen-international", "dashscope-intl",
|
||||||
"qwen-us", "dashscope-us", "mistral", "avian", "longcat", "modelscope", "novita",
|
"qwen-us", "dashscope-us", "mistral", "avian", "longcat", "modelscope", "novita",
|
||||||
"coding-plan", "alibaba-coding", "qwen-coding", "mimo":
|
"coding-plan", "alibaba-coding", "qwen-coding", "mimo":
|
||||||
// All other OpenAI-compatible HTTP providers
|
// All other OpenAI-compatible HTTP providers
|
||||||
if cfg.APIKey() == "" && cfg.APIBase == "" {
|
if cfg.APIKey() == "" && cfg.APIBase == "" && !isEmptyAPIKeyAllowed(protocol) {
|
||||||
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
|
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
|
||||||
}
|
}
|
||||||
apiBase := cfg.APIBase
|
apiBase := cfg.APIBase
|
||||||
|
|
@ -294,64 +336,30 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func isEmptyAPIKeyAllowed(protocol string) bool {
|
||||||
|
meta, ok := protocolMetaByName[protocol]
|
||||||
|
return ok && meta.emptyAPIKeyAllowed
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsEmptyAPIKeyAllowedForProtocol reports whether a protocol allows requests
|
||||||
|
// without api_key when using its default local endpoint.
|
||||||
|
func IsEmptyAPIKeyAllowedForProtocol(protocol string) bool {
|
||||||
|
protocol = strings.ToLower(strings.TrimSpace(protocol))
|
||||||
|
return isEmptyAPIKeyAllowed(protocol)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultAPIBaseForProtocol returns the configured default API base for a protocol.
|
||||||
|
// It returns empty string if the protocol has no default base.
|
||||||
|
func DefaultAPIBaseForProtocol(protocol string) string {
|
||||||
|
protocol = strings.ToLower(strings.TrimSpace(protocol))
|
||||||
|
return getDefaultAPIBase(protocol)
|
||||||
|
}
|
||||||
|
|
||||||
// getDefaultAPIBase returns the default API base URL for a given protocol.
|
// getDefaultAPIBase returns the default API base URL for a given protocol.
|
||||||
func getDefaultAPIBase(protocol string) string {
|
func getDefaultAPIBase(protocol string) string {
|
||||||
switch protocol {
|
meta, ok := protocolMetaByName[protocol]
|
||||||
case "openai":
|
if !ok {
|
||||||
return "https://api.openai.com/v1"
|
|
||||||
case "openrouter":
|
|
||||||
return "https://openrouter.ai/api/v1"
|
|
||||||
case "litellm":
|
|
||||||
return "http://localhost:4000/v1"
|
|
||||||
case "novita":
|
|
||||||
return "https://api.novita.ai/openai"
|
|
||||||
case "groq":
|
|
||||||
return "https://api.groq.com/openai/v1"
|
|
||||||
case "zhipu":
|
|
||||||
return "https://open.bigmodel.cn/api/paas/v4"
|
|
||||||
case "gemini":
|
|
||||||
return "https://generativelanguage.googleapis.com/v1beta"
|
|
||||||
case "nvidia":
|
|
||||||
return "https://integrate.api.nvidia.com/v1"
|
|
||||||
case "ollama":
|
|
||||||
return "http://localhost:11434/v1"
|
|
||||||
case "moonshot":
|
|
||||||
return "https://api.moonshot.cn/v1"
|
|
||||||
case "shengsuanyun":
|
|
||||||
return "https://router.shengsuanyun.com/api/v1"
|
|
||||||
case "deepseek":
|
|
||||||
return "https://api.deepseek.com/v1"
|
|
||||||
case "cerebras":
|
|
||||||
return "https://api.cerebras.ai/v1"
|
|
||||||
case "vivgrid":
|
|
||||||
return "https://api.vivgrid.com/v1"
|
|
||||||
case "volcengine":
|
|
||||||
return "https://ark.cn-beijing.volces.com/api/v3"
|
|
||||||
case "qwen":
|
|
||||||
return "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
|
||||||
case "qwen-intl", "qwen-international", "dashscope-intl":
|
|
||||||
return "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"
|
|
||||||
case "qwen-us", "dashscope-us":
|
|
||||||
return "https://dashscope-us.aliyuncs.com/compatible-mode/v1"
|
|
||||||
case "coding-plan", "alibaba-coding", "qwen-coding":
|
|
||||||
return "https://coding-intl.dashscope.aliyuncs.com/v1"
|
|
||||||
case "coding-plan-anthropic", "alibaba-coding-anthropic":
|
|
||||||
return "https://coding-intl.dashscope.aliyuncs.com/apps/anthropic"
|
|
||||||
case "vllm":
|
|
||||||
return "http://localhost:8000/v1"
|
|
||||||
case "mistral":
|
|
||||||
return "https://api.mistral.ai/v1"
|
|
||||||
case "avian":
|
|
||||||
return "https://api.avian.io/v1"
|
|
||||||
case "minimax":
|
|
||||||
return "https://api.minimaxi.com/v1"
|
|
||||||
case "longcat":
|
|
||||||
return "https://api.longcat.chat/openai"
|
|
||||||
case "modelscope":
|
|
||||||
return "https://api-inference.modelscope.cn/v1"
|
|
||||||
case "mimo":
|
|
||||||
return "https://api.xiaomimimo.com/v1"
|
|
||||||
default:
|
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
return meta.defaultAPIBase
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -121,6 +121,7 @@ func TestCreateProviderFromConfig_DefaultAPIBase(t *testing.T) {
|
||||||
{"vllm", "vllm"},
|
{"vllm", "vllm"},
|
||||||
{"deepseek", "deepseek"},
|
{"deepseek", "deepseek"},
|
||||||
{"ollama", "ollama"},
|
{"ollama", "ollama"},
|
||||||
|
{"lmstudio", "lmstudio"},
|
||||||
{"longcat", "longcat"},
|
{"longcat", "longcat"},
|
||||||
{"modelscope", "modelscope"},
|
{"modelscope", "modelscope"},
|
||||||
{"mimo", "mimo"},
|
{"mimo", "mimo"},
|
||||||
|
|
@ -153,6 +154,12 @@ func TestGetDefaultAPIBase_LiteLLM(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestGetDefaultAPIBase_LMStudio(t *testing.T) {
|
||||||
|
if got := getDefaultAPIBase("lmstudio"); got != "http://localhost:1234/v1" {
|
||||||
|
t.Fatalf("getDefaultAPIBase(%q) = %q, want %q", "lmstudio", got, "http://localhost:1234/v1")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestCreateProviderFromConfig_LiteLLM(t *testing.T) {
|
func TestCreateProviderFromConfig_LiteLLM(t *testing.T) {
|
||||||
cfg := &config.ModelConfig{
|
cfg := &config.ModelConfig{
|
||||||
ModelName: "test-litellm",
|
ModelName: "test-litellm",
|
||||||
|
|
@ -173,6 +180,85 @@ func TestCreateProviderFromConfig_LiteLLM(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderFromConfig_LocalProviders(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
modelName string
|
||||||
|
model string
|
||||||
|
apiKey string
|
||||||
|
wantModelID string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "LMStudio with API key",
|
||||||
|
modelName: "test-lmstudio",
|
||||||
|
model: "lmstudio/openai/gpt-oss-20b",
|
||||||
|
apiKey: "test-key",
|
||||||
|
wantModelID: "openai/gpt-oss-20b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "LMStudio without API key",
|
||||||
|
modelName: "test-lmstudio",
|
||||||
|
model: "lmstudio/openai/gpt-oss-20b",
|
||||||
|
apiKey: "",
|
||||||
|
wantModelID: "openai/gpt-oss-20b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Ollama with API key",
|
||||||
|
modelName: "test-ollama",
|
||||||
|
model: "ollama/llama3.1:8b",
|
||||||
|
apiKey: "test-key",
|
||||||
|
wantModelID: "llama3.1:8b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Ollama without API key",
|
||||||
|
modelName: "test-ollama",
|
||||||
|
model: "ollama/llama3.1:8b",
|
||||||
|
apiKey: "",
|
||||||
|
wantModelID: "llama3.1:8b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "VLLM with API key",
|
||||||
|
modelName: "test-vllm",
|
||||||
|
model: "vllm/Qwen/Qwen3-8B",
|
||||||
|
apiKey: "test-key",
|
||||||
|
wantModelID: "Qwen/Qwen3-8B",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "VLLM without API key",
|
||||||
|
modelName: "test-vllm",
|
||||||
|
model: "vllm/Qwen/Qwen3-8B",
|
||||||
|
apiKey: "",
|
||||||
|
wantModelID: "Qwen/Qwen3-8B",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
cfg := &config.ModelConfig{
|
||||||
|
ModelName: tt.modelName,
|
||||||
|
Model: tt.model,
|
||||||
|
}
|
||||||
|
if tt.apiKey != "" {
|
||||||
|
cfg.SetAPIKey(tt.apiKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
if provider == nil {
|
||||||
|
t.Fatal("CreateProviderFromConfig() returned nil provider")
|
||||||
|
}
|
||||||
|
if modelID != tt.wantModelID {
|
||||||
|
t.Errorf("modelID = %q, want %q", modelID, tt.wantModelID)
|
||||||
|
}
|
||||||
|
if _, ok := provider.(*HTTPProvider); !ok {
|
||||||
|
t.Fatalf("expected *HTTPProvider, got %T", provider)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestCreateProviderFromConfig_LongCat(t *testing.T) {
|
func TestCreateProviderFromConfig_LongCat(t *testing.T) {
|
||||||
cfg := &config.ModelConfig{
|
cfg := &config.ModelConfig{
|
||||||
ModelName: "test-longcat",
|
ModelName: "test-longcat",
|
||||||
|
|
|
||||||
|
|
@ -41,9 +41,9 @@ func NewGitHubCopilotProvider(uri string, connectMode string, model string) (*Gi
|
||||||
}
|
}
|
||||||
|
|
||||||
session, err := client.CreateSession(context.Background(), &copilot.SessionConfig{
|
session, err := client.CreateSession(context.Background(), &copilot.SessionConfig{
|
||||||
Model: model,
|
Model: model,
|
||||||
OnPermissionRequest: copilot.PermissionHandler.ApproveAll,
|
OnPermissionRequest: copilot.PermissionHandler.ApproveAll,
|
||||||
Hooks: &copilot.SessionHooks{},
|
Hooks: &copilot.SessionHooks{},
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
client.Stop()
|
client.Stop()
|
||||||
|
|
|
||||||
|
|
@ -42,6 +42,23 @@ type Option func(*Provider)
|
||||||
|
|
||||||
const defaultRequestTimeout = common.DefaultRequestTimeout
|
const defaultRequestTimeout = common.DefaultRequestTimeout
|
||||||
|
|
||||||
|
var stripModelPrefixProviders = map[string]struct{}{
|
||||||
|
"litellm": {},
|
||||||
|
"moonshot": {},
|
||||||
|
"nvidia": {},
|
||||||
|
"groq": {},
|
||||||
|
"ollama": {},
|
||||||
|
"deepseek": {},
|
||||||
|
"google": {},
|
||||||
|
"openrouter": {},
|
||||||
|
"zhipu": {},
|
||||||
|
"mistral": {},
|
||||||
|
"vivgrid": {},
|
||||||
|
"minimax": {},
|
||||||
|
"novita": {},
|
||||||
|
"lmstudio": {},
|
||||||
|
}
|
||||||
|
|
||||||
func WithMaxTokensField(maxTokensField string) Option {
|
func WithMaxTokensField(maxTokensField string) Option {
|
||||||
return func(p *Provider) {
|
return func(p *Provider) {
|
||||||
p.maxTokensField = maxTokensField
|
p.maxTokensField = maxTokensField
|
||||||
|
|
@ -397,13 +414,11 @@ func normalizeModel(model, apiBase string) string {
|
||||||
}
|
}
|
||||||
|
|
||||||
prefix := strings.ToLower(before)
|
prefix := strings.ToLower(before)
|
||||||
switch prefix {
|
if _, ok := stripModelPrefixProviders[prefix]; ok {
|
||||||
case "litellm", "moonshot", "nvidia", "groq", "ollama", "deepseek", "google",
|
|
||||||
"openrouter", "zhipu", "mistral", "vivgrid", "minimax", "novita":
|
|
||||||
return after
|
return after
|
||||||
default:
|
|
||||||
return model
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return model
|
||||||
}
|
}
|
||||||
|
|
||||||
func buildToolsList(tools []ToolDefinition, nativeSearch bool) []any {
|
func buildToolsList(tools []ToolDefinition, nativeSearch bool) []any {
|
||||||
|
|
|
||||||
|
|
@ -432,7 +432,7 @@ func TestProviderChat_StripsMoonshotPrefixAndNormalizesKimiTemperature(t *testin
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestProviderChat_StripsGroqOllamaDeepseekVivgridNovitaPrefixes(t *testing.T) {
|
func TestProviderChat_StripsKnownProviderPrefixes(t *testing.T) {
|
||||||
var requestBody map[string]any
|
var requestBody map[string]any
|
||||||
|
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
|
@ -474,6 +474,11 @@ func TestProviderChat_StripsGroqOllamaDeepseekVivgridNovitaPrefixes(t *testing.T
|
||||||
input: "ollama/qwen2.5:14b",
|
input: "ollama/qwen2.5:14b",
|
||||||
wantModel: "qwen2.5:14b",
|
wantModel: "qwen2.5:14b",
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
name: "strips lmstudio prefix and keeps nested model",
|
||||||
|
input: "lmstudio/openai/gpt-oss-20b",
|
||||||
|
wantModel: "openai/gpt-oss-20b",
|
||||||
|
},
|
||||||
{
|
{
|
||||||
name: "strips deepseek prefix",
|
name: "strips deepseek prefix",
|
||||||
input: "deepseek/deepseek-chat",
|
input: "deepseek/deepseek-chat",
|
||||||
|
|
@ -579,6 +584,9 @@ func TestNormalizeModel_UsesAPIBase(t *testing.T) {
|
||||||
if got := normalizeModel("deepseek/deepseek-chat", "https://api.deepseek.com/v1"); got != "deepseek-chat" {
|
if got := normalizeModel("deepseek/deepseek-chat", "https://api.deepseek.com/v1"); got != "deepseek-chat" {
|
||||||
t.Fatalf("normalizeModel(deepseek) = %q, want %q", got, "deepseek-chat")
|
t.Fatalf("normalizeModel(deepseek) = %q, want %q", got, "deepseek-chat")
|
||||||
}
|
}
|
||||||
|
if got := normalizeModel("lmstudio/openai/gpt-oss-20b", "http://localhost:1234/v1"); got != "openai/gpt-oss-20b" {
|
||||||
|
t.Fatalf("normalizeModel(lmstudio) = %q, want %q", got, "openai/gpt-oss-20b")
|
||||||
|
}
|
||||||
if got := normalizeModel("openrouter/auto", "https://openrouter.ai/api/v1"); got != "openrouter/auto" {
|
if got := normalizeModel("openrouter/auto", "https://openrouter.ai/api/v1"); got != "openrouter/auto" {
|
||||||
t.Fatalf("normalizeModel(openrouter) = %q, want %q", got, "openrouter/auto")
|
t.Fatalf("normalizeModel(openrouter) = %q, want %q", got, "openrouter/auto")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -59,6 +59,16 @@ func refreshPicoTokensLocked(configPath string) {
|
||||||
gateway.picoToken = cfg.Channels.Pico.Token.String()
|
gateway.picoToken = cfg.Channels.Pico.Token.String()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ensurePicoTokenCachedLocked lazily fills the in-memory pico token cache when
|
||||||
|
// the launcher has already discovered a running gateway via pidData, but has
|
||||||
|
// not yet refreshed the token into memory.
|
||||||
|
func ensurePicoTokenCachedLocked(configPath string) {
|
||||||
|
if gateway.picoToken != "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
refreshPicoTokensLocked(configPath)
|
||||||
|
}
|
||||||
|
|
||||||
const (
|
const (
|
||||||
protocolKey = "Sec-Websocket-Protocol"
|
protocolKey = "Sec-Websocket-Protocol"
|
||||||
tokenPrefix = "token."
|
tokenPrefix = "token."
|
||||||
|
|
|
||||||
|
|
@ -190,12 +190,20 @@ func joinClientVisibleHostPort(r *http.Request, host string, serverListenPort in
|
||||||
func (h *Handler) picoWebUIAddr(r *http.Request) string {
|
func (h *Handler) picoWebUIAddr(r *http.Request) string {
|
||||||
wsPort := h.serverPort
|
wsPort := h.serverPort
|
||||||
if wsPort == 0 {
|
if wsPort == 0 {
|
||||||
wsPort = 18800 // default web server port
|
wsPort = 18800
|
||||||
}
|
}
|
||||||
if fwdHost := forwardedHostFirst(r); fwdHost != "" {
|
if fwdHost := forwardedHostFirst(r); fwdHost != "" {
|
||||||
return joinClientVisibleHostPort(r, fwdHost, wsPort)
|
return joinClientVisibleHostPort(r, fwdHost, wsPort)
|
||||||
}
|
}
|
||||||
host := requestHostName(r)
|
host := requestHostName(r)
|
||||||
|
// Use clientVisiblePort only when an explicit port is present in headers
|
||||||
|
// or Host header — do not infer from TLS/scheme, as serverPort takes priority.
|
||||||
|
if p := forwardedPortFirst(r); p != "" {
|
||||||
|
return net.JoinHostPort(host, p)
|
||||||
|
}
|
||||||
|
if _, port, err := net.SplitHostPort(r.Host); err == nil && port != "" {
|
||||||
|
return net.JoinHostPort(host, port)
|
||||||
|
}
|
||||||
return net.JoinHostPort(host, strconv.Itoa(wsPort))
|
return net.JoinHostPort(host, strconv.Itoa(wsPort))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -68,6 +68,7 @@ func resetGatewayTestState(t *testing.T) {
|
||||||
originalRestartGracePeriod := gatewayRestartGracePeriod
|
originalRestartGracePeriod := gatewayRestartGracePeriod
|
||||||
originalRestartForceKillWindow := gatewayRestartForceKillWindow
|
originalRestartForceKillWindow := gatewayRestartForceKillWindow
|
||||||
originalRestartPollInterval := gatewayRestartPollInterval
|
originalRestartPollInterval := gatewayRestartPollInterval
|
||||||
|
t.Setenv("PICOCLAW_HOME", t.TempDir())
|
||||||
t.Cleanup(func() {
|
t.Cleanup(func() {
|
||||||
gatewayHealthGet = originalHealthGet
|
gatewayHealthGet = originalHealthGet
|
||||||
gatewayRestartGracePeriod = originalRestartGracePeriod
|
gatewayRestartGracePeriod = originalRestartGracePeriod
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,7 @@ import (
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
)
|
)
|
||||||
|
|
||||||
const modelProbeTimeout = 800 * time.Millisecond
|
const modelProbeTimeout = 800 * time.Millisecond
|
||||||
|
|
@ -60,10 +61,14 @@ func requiresRuntimeProbe(m *config.ModelConfig) bool {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
switch modelProtocol(m.Model) {
|
protocol := modelProtocol(m.Model)
|
||||||
|
|
||||||
|
switch protocol {
|
||||||
case "claude-cli", "claudecli", "codex-cli", "codexcli", "github-copilot", "copilot":
|
case "claude-cli", "claudecli", "codex-cli", "codexcli", "github-copilot", "copilot":
|
||||||
return true
|
return true
|
||||||
case "ollama", "vllm":
|
}
|
||||||
|
|
||||||
|
if providers.IsEmptyAPIKeyAllowedForProtocol(protocol) {
|
||||||
apiBase := strings.TrimSpace(m.APIBase)
|
apiBase := strings.TrimSpace(m.APIBase)
|
||||||
return apiBase == "" || hasLocalAPIBase(apiBase)
|
return apiBase == "" || hasLocalAPIBase(apiBase)
|
||||||
}
|
}
|
||||||
|
|
@ -81,7 +86,7 @@ func probeLocalModelAvailability(m *config.ModelConfig) bool {
|
||||||
switch protocol {
|
switch protocol {
|
||||||
case "ollama":
|
case "ollama":
|
||||||
return probeOllamaModelFunc(apiBase, modelID)
|
return probeOllamaModelFunc(apiBase, modelID)
|
||||||
case "vllm":
|
case "vllm", "lmstudio":
|
||||||
return probeOpenAICompatibleModelFunc(apiBase, modelID, m.APIKey())
|
return probeOpenAICompatibleModelFunc(apiBase, modelID, m.APIKey())
|
||||||
case "github-copilot", "copilot":
|
case "github-copilot", "copilot":
|
||||||
return probeTCPServiceFunc(apiBase)
|
return probeTCPServiceFunc(apiBase)
|
||||||
|
|
@ -100,11 +105,12 @@ func modelProbeAPIBase(m *config.ModelConfig) string {
|
||||||
return normalizeModelProbeAPIBase(apiBase)
|
return normalizeModelProbeAPIBase(apiBase)
|
||||||
}
|
}
|
||||||
|
|
||||||
switch modelProtocol(m.Model) {
|
protocol := modelProtocol(m.Model)
|
||||||
case "ollama":
|
if providers.IsEmptyAPIKeyAllowedForProtocol(protocol) {
|
||||||
return "http://localhost:11434/v1"
|
return providers.DefaultAPIBaseForProtocol(protocol)
|
||||||
case "vllm":
|
}
|
||||||
return "http://localhost:8000/v1"
|
|
||||||
|
switch protocol {
|
||||||
case "github-copilot", "copilot":
|
case "github-copilot", "copilot":
|
||||||
return "localhost:4321"
|
return "localhost:4321"
|
||||||
default:
|
default:
|
||||||
|
|
|
||||||
|
|
@ -35,3 +35,53 @@ func TestProbeLocalModelAvailability_OpenAICompatibleIncludesAPIKey(t *testing.T
|
||||||
t.Fatal("probeLocalModelAvailability() = false, want true when api_key is configured")
|
t.Fatal("probeLocalModelAvailability() = false, want true when api_key is configured")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRequiresRuntimeProbe_LMStudio(t *testing.T) {
|
||||||
|
if !requiresRuntimeProbe(&config.ModelConfig{
|
||||||
|
Model: "lmstudio/openai/gpt-oss-20b",
|
||||||
|
}) {
|
||||||
|
t.Fatal("requiresRuntimeProbe(lmstudio with default base) = false, want true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if requiresRuntimeProbe(&config.ModelConfig{
|
||||||
|
Model: "lmstudio/openai/gpt-oss-20b",
|
||||||
|
APIBase: "https://api.example.com/v1",
|
||||||
|
}) {
|
||||||
|
t.Fatal("requiresRuntimeProbe(lmstudio with remote base) = true, want false")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestModelProbeAPIBase_LMStudioDefault(t *testing.T) {
|
||||||
|
got := modelProbeAPIBase(&config.ModelConfig{Model: "lmstudio/openai/gpt-oss-20b"})
|
||||||
|
if got != "http://localhost:1234/v1" {
|
||||||
|
t.Fatalf("modelProbeAPIBase(lmstudio) = %q, want %q", got, "http://localhost:1234/v1")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProbeLocalModelAvailability_LMStudioUsesOpenAICompatibleProbe(t *testing.T) {
|
||||||
|
originalProbe := probeOpenAICompatibleModelFunc
|
||||||
|
defer func() { probeOpenAICompatibleModelFunc = originalProbe }()
|
||||||
|
|
||||||
|
called := false
|
||||||
|
probeOpenAICompatibleModelFunc = func(apiBase, modelID, apiKey string) bool {
|
||||||
|
called = true
|
||||||
|
if apiBase != "http://localhost:1234/v1" {
|
||||||
|
t.Fatalf("apiBase = %q, want %q", apiBase, "http://localhost:1234/v1")
|
||||||
|
}
|
||||||
|
if modelID != "openai/gpt-oss-20b" {
|
||||||
|
t.Fatalf("modelID = %q, want %q", modelID, "openai/gpt-oss-20b")
|
||||||
|
}
|
||||||
|
if apiKey != "" {
|
||||||
|
t.Fatalf("apiKey = %q, want empty", apiKey)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
model := &config.ModelConfig{Model: "lmstudio/openai/gpt-oss-20b"}
|
||||||
|
if !probeLocalModelAvailability(model) {
|
||||||
|
t.Fatal("probeLocalModelAvailability(lmstudio) = false, want true")
|
||||||
|
}
|
||||||
|
if !called {
|
||||||
|
t.Fatal("probeOpenAICompatibleModelFunc was not called for lmstudio")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -56,6 +56,7 @@ func (h *Handler) createWsProxy(origProtocol string, token string) *httputil.Rev
|
||||||
func (h *Handler) handleWebSocketProxy() http.HandlerFunc {
|
func (h *Handler) handleWebSocketProxy() http.HandlerFunc {
|
||||||
return func(w http.ResponseWriter, r *http.Request) {
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
gateway.mu.Lock()
|
gateway.mu.Lock()
|
||||||
|
ensurePicoTokenCachedLocked(h.configPath)
|
||||||
gatewayAvailable := gateway.pidData != nil
|
gatewayAvailable := gateway.pidData != nil
|
||||||
gateway.mu.Unlock()
|
gateway.mu.Unlock()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -377,6 +377,55 @@ func TestHandleWebSocketProxyReloadsGatewayTargetFromConfig(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestHandleWebSocketProxyLoadsCachedPicoTokenWhenMissing(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "config.json")
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
handler := h.handleWebSocketProxy()
|
||||||
|
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.URL.Path != "/pico/ws" {
|
||||||
|
t.Fatalf("path = %q, want %q", r.URL.Path, "/pico/ws")
|
||||||
|
}
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
_, _ = io.WriteString(w, "proxied")
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Gateway.Host = "127.0.0.1"
|
||||||
|
cfg.Gateway.Port = mustGatewayTestPort(t, server.URL)
|
||||||
|
cfg.Channels.Pico.Enabled = true
|
||||||
|
cfg.Channels.Pico.SetToken("cached-token")
|
||||||
|
if err := config.SaveConfig(configPath, cfg); err != nil {
|
||||||
|
t.Fatalf("SaveConfig() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
origPidData := gateway.pidData
|
||||||
|
origPicoToken := gateway.picoToken
|
||||||
|
t.Cleanup(func() {
|
||||||
|
gateway.pidData = origPidData
|
||||||
|
gateway.picoToken = origPicoToken
|
||||||
|
})
|
||||||
|
|
||||||
|
gateway.pidData = &ppid.PidFileData{}
|
||||||
|
gateway.picoToken = ""
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/pico/ws?session_id=test-session", nil)
|
||||||
|
req.Header.Set(protocolKey, tokenPrefix+"cached-token")
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
handler(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d", rec.Code, http.StatusOK)
|
||||||
|
}
|
||||||
|
if body := rec.Body.String(); body != "proxied" {
|
||||||
|
t.Fatalf("body = %q, want %q", body, "proxied")
|
||||||
|
}
|
||||||
|
if gateway.picoToken != "cached-token" {
|
||||||
|
t.Fatalf("gateway.picoToken = %q, want %q", gateway.picoToken, "cached-token")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func mustGatewayTestPort(t *testing.T, rawURL string) int {
|
func mustGatewayTestPort(t *testing.T, rawURL string) int {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue