refactor: adopt upstream for add/add conflict files, extract fork code to _ext.go
Replace qq.go, cron.go, telegram.go, session_key.go, interfaces.go, common.go, memory.go, tool_call_extract.go, string.go with upstream versions. Fork-only code extracted to corresponding _ext.go files. Add ReplyToMessageID to bus.OutboundMessage for upstream telegram compat. Reduces conflict markers from 313→253 (-19%), conflict files 34→28 (-18%). Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
ab12296848
commit
431c96186a
18 changed files with 1707 additions and 1767 deletions
3
go.mod
3
go.mod
|
|
@ -10,6 +10,7 @@ require (
|
||||||
github.com/chzyer/readline v1.5.1
|
github.com/chzyer/readline v1.5.1
|
||||||
github.com/ergochat/irc-go v0.5.0
|
github.com/ergochat/irc-go v0.5.0
|
||||||
github.com/gdamore/tcell/v2 v2.13.8
|
github.com/gdamore/tcell/v2 v2.13.8
|
||||||
|
github.com/gomarkdown/markdown v0.0.0-20260217112301-37c66b85d6ab
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/gorilla/websocket v1.5.3
|
github.com/gorilla/websocket v1.5.3
|
||||||
github.com/h2non/filetype v1.1.3
|
github.com/h2non/filetype v1.1.3
|
||||||
|
|
@ -28,6 +29,7 @@ require (
|
||||||
golang.org/x/oauth2 v0.35.0
|
golang.org/x/oauth2 v0.35.0
|
||||||
golang.org/x/time v0.14.0
|
golang.org/x/time v0.14.0
|
||||||
google.golang.org/protobuf v1.36.11
|
google.golang.org/protobuf v1.36.11
|
||||||
|
gopkg.in/yaml.v3 v3.0.1
|
||||||
maunium.net/go/mautrix v0.26.3
|
maunium.net/go/mautrix v0.26.3
|
||||||
modernc.org/sqlite v1.46.1
|
modernc.org/sqlite v1.46.1
|
||||||
)
|
)
|
||||||
|
|
@ -59,7 +61,6 @@ require (
|
||||||
golang.org/x/exp v0.0.0-20260212183809-81e46e3db34a // indirect
|
golang.org/x/exp v0.0.0-20260212183809-81e46e3db34a // indirect
|
||||||
golang.org/x/term v0.40.0 // indirect
|
golang.org/x/term v0.40.0 // indirect
|
||||||
golang.org/x/text v0.34.0 // indirect
|
golang.org/x/text v0.34.0 // indirect
|
||||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
|
||||||
modernc.org/libc v1.67.6 // indirect
|
modernc.org/libc v1.67.6 // indirect
|
||||||
modernc.org/mathutil v1.7.1 // indirect
|
modernc.org/mathutil v1.7.1 // indirect
|
||||||
modernc.org/memory v1.11.0 // indirect
|
modernc.org/memory v1.11.0 // indirect
|
||||||
|
|
|
||||||
4
go.sum
4
go.sum
|
|
@ -79,6 +79,8 @@ github.com/golang/protobuf v1.4.0/go.mod h1:jodUvKwWbYaEsadDk5Fwe5c77LiNKVO9IDvq
|
||||||
github.com/golang/protobuf v1.4.2/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI=
|
github.com/golang/protobuf v1.4.2/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI=
|
||||||
github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk=
|
github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk=
|
||||||
github.com/golang/protobuf v1.5.2/go.mod h1:XVQd3VNwM+JqD3oG2Ue2ip4fOMUkwXdXDdiuN0vRsmY=
|
github.com/golang/protobuf v1.5.2/go.mod h1:XVQd3VNwM+JqD3oG2Ue2ip4fOMUkwXdXDdiuN0vRsmY=
|
||||||
|
github.com/gomarkdown/markdown v0.0.0-20260217112301-37c66b85d6ab h1:VYNivV7P8IRHUam2swVUNkhIdp0LRRFKe4hXNnoZKTc=
|
||||||
|
github.com/gomarkdown/markdown v0.0.0-20260217112301-37c66b85d6ab/go.mod h1:JDGcbDT52eL4fju3sZ4TeHGsQwhG9nbDV21aMyhwPoA=
|
||||||
github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
|
github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
|
||||||
github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
|
github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
|
||||||
github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
||||||
|
|
@ -269,8 +271,6 @@ golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug
|
||||||
golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
|
golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
|
||||||
golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg=
|
golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg=
|
||||||
golang.org/x/net v0.19.0/go.mod h1:CfAk/cbD4CthTvqiEl8NpboMuiuOYsAr/7NOjZJtv1U=
|
golang.org/x/net v0.19.0/go.mod h1:CfAk/cbD4CthTvqiEl8NpboMuiuOYsAr/7NOjZJtv1U=
|
||||||
golang.org/x/net v0.50.0 h1:ucWh9eiCGyDR3vtzso0WMQinm2Dnt8cFMuQa9K33J60=
|
|
||||||
golang.org/x/net v0.50.0/go.mod h1:UgoSli3F/pBgdJBHCTc+tp3gmrU4XswgGRgtnwWTfyM=
|
|
||||||
golang.org/x/net v0.51.0 h1:94R/GTO7mt3/4wIKpcR5gkGmRLOuE/2hNGeWq/GBIFo=
|
golang.org/x/net v0.51.0 h1:94R/GTO7mt3/4wIKpcR5gkGmRLOuE/2hNGeWq/GBIFo=
|
||||||
golang.org/x/net v0.51.0/go.mod h1:aamm+2QF5ogm02fjy5Bb7CQ0WMt1/WVM7FtyaTLlA9Y=
|
golang.org/x/net v0.51.0/go.mod h1:aamm+2QF5ogm02fjy5Bb7CQ0WMt1/WVM7FtyaTLlA9Y=
|
||||||
golang.org/x/oauth2 v0.23.0/go.mod h1:XYTD2NtWslqkgxebSiOHnXEap4TF09sJSc7H1sXbhtI=
|
golang.org/x/oauth2 v0.23.0/go.mod h1:XYTD2NtWslqkgxebSiOHnXEap4TF09sJSc7H1sXbhtI=
|
||||||
|
|
|
||||||
1125
pkg/agent/memory.go
1125
pkg/agent/memory.go
File diff suppressed because it is too large
Load diff
718
pkg/agent/memory_ext.go
Normal file
718
pkg/agent/memory_ext.go
Normal file
|
|
@ -0,0 +1,718 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"regexp"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Cache types for MemoryStore.
|
||||||
|
|
||||||
|
type longTermFileCache struct {
|
||||||
|
loaded bool
|
||||||
|
exists bool
|
||||||
|
modTime time.Time
|
||||||
|
size int64
|
||||||
|
content string
|
||||||
|
}
|
||||||
|
|
||||||
|
type parsedPlanStateCache struct {
|
||||||
|
loaded bool
|
||||||
|
sourceContent string
|
||||||
|
state parsedPlanState
|
||||||
|
}
|
||||||
|
|
||||||
|
type parsedPlanState struct {
|
||||||
|
content string
|
||||||
|
hasActivePlan bool
|
||||||
|
status string
|
||||||
|
currentPhase int
|
||||||
|
totalPhases int
|
||||||
|
workDir string
|
||||||
|
taskName string
|
||||||
|
phases []PlanPhase
|
||||||
|
}
|
||||||
|
|
||||||
|
// InvalidateCache clears all in-memory caches for MEMORY.md content and parsed plan state.
|
||||||
|
func (ms *MemoryStore) InvalidateCache() {
|
||||||
|
ms.cacheMu.Lock()
|
||||||
|
defer ms.cacheMu.Unlock()
|
||||||
|
ms.longTermCache = longTermFileCache{}
|
||||||
|
ms.parsedPlanCache = parsedPlanStateCache{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) readLongTermCached() string {
|
||||||
|
info, err := os.Stat(ms.memoryFile)
|
||||||
|
if err != nil {
|
||||||
|
if !os.IsNotExist(err) {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
ms.cacheMu.RLock()
|
||||||
|
cachedMissing := ms.longTermCache.loaded && !ms.longTermCache.exists
|
||||||
|
ms.cacheMu.RUnlock()
|
||||||
|
if cachedMissing {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
ms.cacheMu.Lock()
|
||||||
|
ms.longTermCache = longTermFileCache{loaded: true, exists: false}
|
||||||
|
ms.parsedPlanCache = parsedPlanStateCache{}
|
||||||
|
ms.cacheMu.Unlock()
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
modTime := info.ModTime()
|
||||||
|
size := info.Size()
|
||||||
|
|
||||||
|
ms.cacheMu.RLock()
|
||||||
|
if ms.longTermCache.loaded &&
|
||||||
|
ms.longTermCache.exists &&
|
||||||
|
ms.longTermCache.modTime.Equal(modTime) &&
|
||||||
|
ms.longTermCache.size == size {
|
||||||
|
content := ms.longTermCache.content
|
||||||
|
ms.cacheMu.RUnlock()
|
||||||
|
return content
|
||||||
|
}
|
||||||
|
ms.cacheMu.RUnlock()
|
||||||
|
|
||||||
|
data, err := os.ReadFile(ms.memoryFile)
|
||||||
|
if err != nil {
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
ms.cacheMu.Lock()
|
||||||
|
ms.longTermCache = longTermFileCache{loaded: true, exists: false}
|
||||||
|
ms.parsedPlanCache = parsedPlanStateCache{}
|
||||||
|
ms.cacheMu.Unlock()
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
content := string(data)
|
||||||
|
ms.cacheMu.Lock()
|
||||||
|
ms.longTermCache = longTermFileCache{
|
||||||
|
loaded: true,
|
||||||
|
exists: true,
|
||||||
|
modTime: modTime,
|
||||||
|
size: size,
|
||||||
|
content: content,
|
||||||
|
}
|
||||||
|
if ms.parsedPlanCache.loaded && ms.parsedPlanCache.sourceContent != content {
|
||||||
|
ms.parsedPlanCache = parsedPlanStateCache{}
|
||||||
|
}
|
||||||
|
ms.cacheMu.Unlock()
|
||||||
|
return content
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) getParsedPlanState() parsedPlanState {
|
||||||
|
content := ms.ReadLongTerm()
|
||||||
|
|
||||||
|
ms.cacheMu.RLock()
|
||||||
|
if ms.parsedPlanCache.loaded && ms.parsedPlanCache.sourceContent == content {
|
||||||
|
state := ms.parsedPlanCache.state
|
||||||
|
ms.cacheMu.RUnlock()
|
||||||
|
return state
|
||||||
|
}
|
||||||
|
ms.cacheMu.RUnlock()
|
||||||
|
|
||||||
|
state := ms.parsePlanState(content)
|
||||||
|
|
||||||
|
ms.cacheMu.Lock()
|
||||||
|
if !ms.parsedPlanCache.loaded || ms.parsedPlanCache.sourceContent != content {
|
||||||
|
ms.parsedPlanCache = parsedPlanStateCache{
|
||||||
|
loaded: true,
|
||||||
|
sourceContent: content,
|
||||||
|
state: state,
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
state = ms.parsedPlanCache.state
|
||||||
|
}
|
||||||
|
ms.cacheMu.Unlock()
|
||||||
|
return state
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) parsePlanState(content string) parsedPlanState {
|
||||||
|
state := parsedPlanState{content: content}
|
||||||
|
if content == "" || !reActivePlan.MatchString(content) {
|
||||||
|
return state
|
||||||
|
}
|
||||||
|
state.hasActivePlan = true
|
||||||
|
if m := reStatus.FindStringSubmatch(content); len(m) >= 2 {
|
||||||
|
state.status = strings.TrimSpace(m[1])
|
||||||
|
}
|
||||||
|
if m := rePhase.FindStringSubmatch(content); len(m) >= 2 {
|
||||||
|
state.currentPhase, _ = strconv.Atoi(m[1])
|
||||||
|
}
|
||||||
|
state.totalPhases = maxPhaseNumber(content)
|
||||||
|
if m := reWorkDir.FindStringSubmatch(content); len(m) >= 2 {
|
||||||
|
state.workDir = strings.TrimSpace(m[1])
|
||||||
|
}
|
||||||
|
if m := reTaskLine.FindStringSubmatch(content); len(m) >= 2 {
|
||||||
|
state.taskName = strings.TrimSpace(m[1])
|
||||||
|
}
|
||||||
|
state.phases = ms.getPlanPhasesFrom(content)
|
||||||
|
return state
|
||||||
|
}
|
||||||
|
|
||||||
|
func clonePlanPhases(phases []PlanPhase) []PlanPhase {
|
||||||
|
if len(phases) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
result := make([]PlanPhase, 0, len(phases))
|
||||||
|
for _, p := range phases {
|
||||||
|
phase := PlanPhase{Number: p.Number, Title: p.Title}
|
||||||
|
if len(p.Steps) > 0 {
|
||||||
|
phase.Steps = append([]PlanStep(nil), p.Steps...)
|
||||||
|
}
|
||||||
|
result = append(result, phase)
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClearLongTerm removes the long-term memory file.
|
||||||
|
func (ms *MemoryStore) ClearLongTerm() error {
|
||||||
|
if err := os.Remove(ms.memoryFile); err != nil && !os.IsNotExist(err) {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
ms.InvalidateCache()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------- Plan state query methods ----------
|
||||||
|
|
||||||
|
var (
|
||||||
|
reActivePlan = regexp.MustCompile(`(?m)^# Active Plan`)
|
||||||
|
reStatus = regexp.MustCompile(`(?m)^> Status:\s*(.+)`)
|
||||||
|
rePhase = regexp.MustCompile(`(?m)^> Phase:\s*(\d+)`)
|
||||||
|
rePhaseHeader = regexp.MustCompile(`(?m)^## Phase (\d+):\s*(.*)`)
|
||||||
|
reWorkDir = regexp.MustCompile(`(?m)^> WorkDir:\s*(.+)`)
|
||||||
|
reTaskLine = regexp.MustCompile(`(?m)^> Task:\s*(.+)`)
|
||||||
|
)
|
||||||
|
|
||||||
|
// HasActivePlan returns true if MEMORY.md contains an active plan.
|
||||||
|
func (ms *MemoryStore) HasActivePlan() bool {
|
||||||
|
return ms.getParsedPlanState().hasActivePlan
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPlanStatus returns the plan status: "interviewing", "executing", or "".
|
||||||
|
func (ms *MemoryStore) GetPlanStatus() string {
|
||||||
|
return ms.getParsedPlanState().status
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetCurrentPhase returns the current phase number from "> Phase: N".
|
||||||
|
func (ms *MemoryStore) GetCurrentPhase() int {
|
||||||
|
return ms.getParsedPlanState().currentPhase
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetTotalPhases returns the total number of phases (max ## Phase N).
|
||||||
|
func (ms *MemoryStore) GetTotalPhases() int {
|
||||||
|
return ms.getParsedPlanState().totalPhases
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsPlanComplete returns true if all steps in all phases are [x].
|
||||||
|
func (ms *MemoryStore) IsPlanComplete() bool {
|
||||||
|
phases := ms.getParsedPlanState().phases
|
||||||
|
if len(phases) == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
hasSteps := false
|
||||||
|
for _, p := range phases {
|
||||||
|
for _, s := range p.Steps {
|
||||||
|
hasSteps = true
|
||||||
|
if !s.Done {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return hasSteps
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsCurrentPhaseComplete returns true if all steps in the current phase are [x].
|
||||||
|
func (ms *MemoryStore) IsCurrentPhaseComplete() bool {
|
||||||
|
state := ms.getParsedPlanState()
|
||||||
|
if state.currentPhase == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, p := range state.phases {
|
||||||
|
if p.Number == state.currentPhase {
|
||||||
|
if len(p.Steps) == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, s := range p.Steps {
|
||||||
|
if !s.Done {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) extractPhaseContent(content string, phase int) string {
|
||||||
|
lines := strings.Split(content, "\n")
|
||||||
|
inPhase := false
|
||||||
|
var result []string
|
||||||
|
phasePrefix := fmt.Sprintf("## Phase %d:", phase)
|
||||||
|
for _, line := range lines {
|
||||||
|
if strings.HasPrefix(line, phasePrefix) {
|
||||||
|
inPhase = true
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if inPhase {
|
||||||
|
if strings.HasPrefix(line, "## Phase ") || strings.HasPrefix(line, "## Context") {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
result = append(result, line)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return strings.Join(result, "\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
// PlanPhase represents a phase with its steps, for structured API output.
|
||||||
|
type PlanPhase struct {
|
||||||
|
Number int `json:"number"`
|
||||||
|
Title string `json:"title"`
|
||||||
|
Steps []PlanStep `json:"steps"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PlanStep represents a single step within a phase.
|
||||||
|
type PlanStep struct {
|
||||||
|
Index int `json:"index"` // 1-based within the phase
|
||||||
|
Description string `json:"description"`
|
||||||
|
Done bool `json:"done"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPlanPhases parses MEMORY.md and returns all phases with their steps.
|
||||||
|
func (ms *MemoryStore) GetPlanPhases() []PlanPhase {
|
||||||
|
return clonePlanPhases(ms.getParsedPlanState().phases)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) getPlanPhasesFrom(content string) []PlanPhase {
|
||||||
|
if !reActivePlan.MatchString(content) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
totalPhases := maxPhaseNumber(content)
|
||||||
|
phases := make([]PlanPhase, 0, totalPhases)
|
||||||
|
for p := 1; p <= totalPhases; p++ {
|
||||||
|
title := ms.getPhaseTitle(content, p)
|
||||||
|
phaseContent := ms.extractPhaseContent(content, p)
|
||||||
|
var steps []PlanStep
|
||||||
|
stepIdx := 0
|
||||||
|
for _, line := range strings.Split(phaseContent, "\n") {
|
||||||
|
line = strings.TrimSpace(line)
|
||||||
|
if strings.HasPrefix(line, "- [x] ") {
|
||||||
|
stepIdx++
|
||||||
|
steps = append(steps, PlanStep{Index: stepIdx, Description: line[6:], Done: true})
|
||||||
|
} else if strings.HasPrefix(line, "- [ ] ") {
|
||||||
|
stepIdx++
|
||||||
|
steps = append(steps, PlanStep{Index: stepIdx, Description: line[6:], Done: false})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
phases = append(phases, PlanPhase{Number: p, Title: title, Steps: steps})
|
||||||
|
}
|
||||||
|
return phases
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------- Plan mutation methods ----------
|
||||||
|
|
||||||
|
// SetStatus sets the plan status (interviewing or executing).
|
||||||
|
func (ms *MemoryStore) SetStatus(status string) error {
|
||||||
|
content := ms.ReadLongTerm()
|
||||||
|
if m := reStatus.FindString(content); m != "" {
|
||||||
|
content = strings.Replace(content, m, "> Status: "+status, 1)
|
||||||
|
}
|
||||||
|
return ms.WriteLongTerm(content)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AdvancePhase increments the current phase number by 1.
|
||||||
|
func (ms *MemoryStore) AdvancePhase() error {
|
||||||
|
content := ms.ReadLongTerm()
|
||||||
|
m := rePhase.FindStringSubmatch(content)
|
||||||
|
if len(m) < 2 {
|
||||||
|
return fmt.Errorf("no phase marker found")
|
||||||
|
}
|
||||||
|
current, _ := strconv.Atoi(m[1])
|
||||||
|
next := current + 1
|
||||||
|
content = strings.Replace(content, m[0], fmt.Sprintf("> Phase: %d", next), 1)
|
||||||
|
return ms.WriteLongTerm(content)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetPhase sets the current phase number to n.
|
||||||
|
func (ms *MemoryStore) SetPhase(n int) error {
|
||||||
|
content := ms.ReadLongTerm()
|
||||||
|
m := rePhase.FindString(content)
|
||||||
|
if m == "" {
|
||||||
|
return fmt.Errorf("no phase marker found")
|
||||||
|
}
|
||||||
|
content = strings.Replace(content, m, fmt.Sprintf("> Phase: %d", n), 1)
|
||||||
|
return ms.WriteLongTerm(content)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarkStep marks the nth step (1-based) in the given phase as done [x].
|
||||||
|
func (ms *MemoryStore) MarkStep(phase, step int) error {
|
||||||
|
content := ms.ReadLongTerm()
|
||||||
|
lines := strings.Split(content, "\n")
|
||||||
|
phasePrefix := fmt.Sprintf("## Phase %d:", phase)
|
||||||
|
inPhase := false
|
||||||
|
stepCount := 0
|
||||||
|
for i, line := range lines {
|
||||||
|
if strings.HasPrefix(line, phasePrefix) {
|
||||||
|
inPhase = true
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if inPhase {
|
||||||
|
if strings.HasPrefix(line, "## Phase ") || strings.HasPrefix(line, "## Context") {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(line, "- [ ] ") {
|
||||||
|
stepCount++
|
||||||
|
if stepCount == step {
|
||||||
|
lines[i] = strings.Replace(line, "- [ ] ", "- [x] ", 1)
|
||||||
|
return ms.WriteLongTerm(strings.Join(lines, "\n"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fmt.Errorf("step %d not found in phase %d", step, phase)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddStep appends a new step to the given phase.
|
||||||
|
func (ms *MemoryStore) AddStep(phase int, desc string) error {
|
||||||
|
content := ms.ReadLongTerm()
|
||||||
|
lines := strings.Split(content, "\n")
|
||||||
|
phasePrefix := fmt.Sprintf("## Phase %d:", phase)
|
||||||
|
inPhase := false
|
||||||
|
insertIdx := -1
|
||||||
|
for i, line := range lines {
|
||||||
|
if strings.HasPrefix(line, phasePrefix) {
|
||||||
|
inPhase = true
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if inPhase {
|
||||||
|
if strings.HasPrefix(line, "## Phase ") || strings.HasPrefix(line, "## Context") {
|
||||||
|
insertIdx = i
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(line, "- [") {
|
||||||
|
insertIdx = i + 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if insertIdx < 0 {
|
||||||
|
if inPhase {
|
||||||
|
insertIdx = len(lines)
|
||||||
|
} else {
|
||||||
|
return fmt.Errorf("phase %d not found", phase)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
newStep := "- [ ] " + desc
|
||||||
|
newLines := make([]string, 0, len(lines)+1)
|
||||||
|
newLines = append(newLines, lines[:insertIdx]...)
|
||||||
|
newLines = append(newLines, newStep)
|
||||||
|
newLines = append(newLines, lines[insertIdx:]...)
|
||||||
|
return ms.WriteLongTerm(strings.Join(newLines, "\n"))
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidatePlanStructure checks that the plan has valid structure for
|
||||||
|
// transitioning out of the interview phase.
|
||||||
|
func (ms *MemoryStore) ValidatePlanStructure() error {
|
||||||
|
content := ms.ReadLongTerm()
|
||||||
|
if !reActivePlan.MatchString(content) {
|
||||||
|
return fmt.Errorf("missing '# Active Plan' header")
|
||||||
|
}
|
||||||
|
if !reStatus.MatchString(content) {
|
||||||
|
return fmt.Errorf("missing '> Status:' line")
|
||||||
|
}
|
||||||
|
if !rePhase.MatchString(content) {
|
||||||
|
return fmt.Errorf("missing '> Phase:' line")
|
||||||
|
}
|
||||||
|
phases := ms.getPlanPhasesFrom(content)
|
||||||
|
if len(phases) == 0 {
|
||||||
|
return fmt.Errorf("no '## Phase N:' sections found")
|
||||||
|
}
|
||||||
|
for _, p := range phases {
|
||||||
|
if len(p.Steps) == 0 {
|
||||||
|
return fmt.Errorf("Phase %d has no checkbox steps (use '- [ ] ...')", p.Number)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------- Selective injection methods ----------
|
||||||
|
|
||||||
|
// GetPlanWorkDir returns the WorkDir from the plan metadata, or "".
|
||||||
|
func (ms *MemoryStore) GetPlanWorkDir() string {
|
||||||
|
return ms.getParsedPlanState().workDir
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPlanTaskName returns the task description from the plan metadata, or "".
|
||||||
|
func (ms *MemoryStore) GetPlanTaskName() string {
|
||||||
|
return ms.getParsedPlanState().taskName
|
||||||
|
}
|
||||||
|
|
||||||
|
const interviewSeedTemplate = `# Active Plan
|
||||||
|
|
||||||
|
> Task: %s
|
||||||
|
> WorkDir: %s
|
||||||
|
> Status: interviewing
|
||||||
|
> Phase: 1
|
||||||
|
`
|
||||||
|
|
||||||
|
// BuildInterviewSeed creates the initial plan seed for a given task description.
|
||||||
|
func BuildInterviewSeed(task, workDir string) string {
|
||||||
|
return fmt.Sprintf(interviewSeedTemplate, task, workDir)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetInterviewContext returns context for injection during the interviewing phase.
|
||||||
|
func (ms *MemoryStore) GetInterviewContext() string {
|
||||||
|
return ms.getInterviewContextFrom(ms.ReadLongTerm())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) getInterviewContextFrom(content string) string {
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("## Active Plan (interviewing)\n\n")
|
||||||
|
sb.WriteString(content)
|
||||||
|
sb.WriteString("\n\n### Interview Guide\n")
|
||||||
|
sb.WriteString("Ask about:\n")
|
||||||
|
sb.WriteString("- Goals and success criteria\n")
|
||||||
|
sb.WriteString("- Constraints (time, budget, platform)\n")
|
||||||
|
sb.WriteString("- Environment (OS, language, runtime versions)\n")
|
||||||
|
sb.WriteString("- Tooling preferences (test framework, linter, formatter, CI)\n")
|
||||||
|
sb.WriteString("- Key commands the user already runs (build, test, deploy)\n")
|
||||||
|
sb.WriteString("\n### Rules\n")
|
||||||
|
sb.WriteString("- NEVER remove or overwrite the header block (`# Active Plan`, `> Task:`, `> Status:`, `> Phase:` lines). The system parses these to track state.\n")
|
||||||
|
sb.WriteString("- After each answer, use edit_file to append findings to the ## Context section of memory/MEMORY.md.\n")
|
||||||
|
sb.WriteString("- When you have enough information, use edit_file to add ## Phase, ## Commands, and ## Context sections BELOW the header block.\n")
|
||||||
|
sb.WriteString("- Each step MUST use checkbox syntax: `- [ ] description`. The system parses checkboxes to track progress.\n")
|
||||||
|
sb.WriteString("- Organize into 2-5 phases with 3-5 steps each.\n")
|
||||||
|
sb.WriteString("- After writing Phases, change `> Status: interviewing` to `> Status: review` via edit_file. The user must approve with /plan start before execution begins.\n")
|
||||||
|
sb.WriteString("\n### Target Format (MANDATORY — system parses this exact structure)\n\n")
|
||||||
|
sb.WriteString("# Active Plan\n")
|
||||||
|
sb.WriteString("> Task: <description>\n")
|
||||||
|
sb.WriteString("> WorkDir: <path>\n")
|
||||||
|
sb.WriteString("> Status: interviewing\n")
|
||||||
|
sb.WriteString("> Phase: 1\n\n")
|
||||||
|
sb.WriteString("## Phase 1: <title>\n")
|
||||||
|
sb.WriteString("- [ ] Step description\n")
|
||||||
|
sb.WriteString("- [ ] Step description\n\n")
|
||||||
|
sb.WriteString("## Phase 2: <title>\n")
|
||||||
|
sb.WriteString("- [ ] Step description\n")
|
||||||
|
sb.WriteString("- [ ] Step description\n\n")
|
||||||
|
sb.WriteString("## Commands\n")
|
||||||
|
sb.WriteString("build: <project-specific build command>\n")
|
||||||
|
sb.WriteString("test: <project-specific test command>\n")
|
||||||
|
sb.WriteString("lint: <project-specific lint command>\n\n")
|
||||||
|
sb.WriteString("## Context\n")
|
||||||
|
sb.WriteString("<collected requirements, decisions, environment>\n")
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetReviewContext returns context for injection during the review phase.
|
||||||
|
func (ms *MemoryStore) GetReviewContext() string {
|
||||||
|
return ms.getReviewContextFrom(ms.ReadLongTerm())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) getReviewContextFrom(content string) string {
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("## Active Plan (awaiting approval)\n\n")
|
||||||
|
sb.WriteString(content)
|
||||||
|
sb.WriteString("\n\nThe plan is awaiting user approval.\n")
|
||||||
|
sb.WriteString("- If the user requests changes, update memory/MEMORY.md via edit_file.\n")
|
||||||
|
sb.WriteString("- Do NOT change Status yourself. The user will run /plan start to approve.\n")
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPlanContext returns context for injection during the executing phase.
|
||||||
|
func (ms *MemoryStore) GetPlanContext() string {
|
||||||
|
return ms.getPlanContextFrom(ms.ReadLongTerm())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) getPlanContextFrom(content string) string {
|
||||||
|
var currentPhase int
|
||||||
|
if m := rePhase.FindStringSubmatch(content); len(m) >= 2 {
|
||||||
|
currentPhase, _ = strconv.Atoi(m[1])
|
||||||
|
}
|
||||||
|
totalPhases := maxPhaseNumber(content)
|
||||||
|
taskLine := ""
|
||||||
|
if m := reTaskLine.FindStringSubmatch(content); len(m) >= 2 {
|
||||||
|
taskLine = strings.TrimSpace(m[1])
|
||||||
|
}
|
||||||
|
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("## Active Plan\n")
|
||||||
|
fmt.Fprintf(&sb, "Task: %s | Phase %d/%d\n", taskLine, currentPhase, totalPhases)
|
||||||
|
|
||||||
|
for p := 1; p < currentPhase; p++ {
|
||||||
|
title := ms.getPhaseTitle(content, p)
|
||||||
|
fmt.Fprintf(&sb, "Done: Phase %d (%s)\n", p, title)
|
||||||
|
}
|
||||||
|
|
||||||
|
if currentPhase > 0 {
|
||||||
|
title := ms.getPhaseTitle(content, currentPhase)
|
||||||
|
fmt.Fprintf(&sb, "### Current: Phase %d — %s\n", currentPhase, title)
|
||||||
|
phaseContent := ms.extractPhaseContent(content, currentPhase)
|
||||||
|
sb.WriteString(strings.TrimSpace(phaseContent))
|
||||||
|
sb.WriteString("\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
if commandsContent := ms.extractCommandsSection(content); commandsContent != "" {
|
||||||
|
sb.WriteString("### Commands\n")
|
||||||
|
sb.WriteString(commandsContent)
|
||||||
|
sb.WriteString("\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
if contextContent := ms.extractContextSection(content); contextContent != "" {
|
||||||
|
sb.WriteString("### Context\n")
|
||||||
|
sb.WriteString(contextContent)
|
||||||
|
sb.WriteString("\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
if orchContent := ms.extractSection(content, "Orchestration"); orchContent != "" {
|
||||||
|
sb.WriteString("### Orchestration\n")
|
||||||
|
sb.WriteString(orchContent)
|
||||||
|
sb.WriteString("\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func maxPhaseNumber(content string) int {
|
||||||
|
matches := rePhaseHeader.FindAllStringSubmatch(content, -1)
|
||||||
|
maxN := 0
|
||||||
|
for _, m := range matches {
|
||||||
|
if len(m) >= 2 {
|
||||||
|
n, _ := strconv.Atoi(m[1])
|
||||||
|
if n > maxN {
|
||||||
|
maxN = n
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return maxN
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) getPhaseTitle(content string, phase int) string {
|
||||||
|
matches := rePhaseHeader.FindAllStringSubmatch(content, -1)
|
||||||
|
for _, m := range matches {
|
||||||
|
if len(m) >= 3 {
|
||||||
|
n, _ := strconv.Atoi(m[1])
|
||||||
|
if n == phase {
|
||||||
|
return strings.TrimSpace(m[2])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) extractSection(content, name string) string {
|
||||||
|
lines := strings.Split(content, "\n")
|
||||||
|
prefix := "## " + name
|
||||||
|
inSection := false
|
||||||
|
var result []string
|
||||||
|
for _, line := range lines {
|
||||||
|
if strings.HasPrefix(line, prefix) {
|
||||||
|
inSection = true
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if inSection {
|
||||||
|
if strings.HasPrefix(line, "## ") {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
result = append(result, line)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return strings.TrimSpace(strings.Join(result, "\n"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) extractContextSection(content string) string {
|
||||||
|
return ms.extractSection(content, "Context")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ms *MemoryStore) extractCommandsSection(content string) string {
|
||||||
|
return ms.extractSection(content, "Commands")
|
||||||
|
}
|
||||||
|
|
||||||
|
// FormatPlanDisplay returns a user-facing display of the full plan with emoji indicators.
|
||||||
|
func (ms *MemoryStore) FormatPlanDisplay() string {
|
||||||
|
state := ms.getParsedPlanState()
|
||||||
|
if !state.hasActivePlan {
|
||||||
|
return "No active plan."
|
||||||
|
}
|
||||||
|
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString(fmt.Sprintf("Plan: %s\n", state.taskName))
|
||||||
|
sb.WriteString(fmt.Sprintf("Status: %s | Phase %d/%d\n\n", state.status, state.currentPhase, len(state.phases)))
|
||||||
|
|
||||||
|
for _, p := range state.phases {
|
||||||
|
var emoji string
|
||||||
|
if p.Number < state.currentPhase {
|
||||||
|
emoji = "\u2705"
|
||||||
|
} else if p.Number == state.currentPhase {
|
||||||
|
emoji = "\u25B6\uFE0F"
|
||||||
|
} else {
|
||||||
|
emoji = "\u23F3"
|
||||||
|
}
|
||||||
|
sb.WriteString(fmt.Sprintf("%s Phase %d: %s\n", emoji, p.Number, p.Title))
|
||||||
|
if p.Number <= state.currentPhase {
|
||||||
|
for _, s := range p.Steps {
|
||||||
|
if s.Done {
|
||||||
|
sb.WriteString(" \u2611 " + s.Description + "\n")
|
||||||
|
} else {
|
||||||
|
sb.WriteString(" \u2610 " + s.Description + "\n")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if commandsContent := ms.extractCommandsSection(state.content); commandsContent != "" {
|
||||||
|
sb.WriteString("\nCommands:\n")
|
||||||
|
for _, line := range strings.Split(commandsContent, "\n") {
|
||||||
|
line = strings.TrimSpace(line)
|
||||||
|
if line != "" {
|
||||||
|
sb.WriteString(" " + line + "\n")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if contextContent := ms.extractContextSection(state.content); contextContent != "" {
|
||||||
|
sb.WriteString("\nContext: " + contextContent + "\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------- GetMemoryContext (plan-aware) ----------
|
||||||
|
|
||||||
|
func (ms *MemoryStore) getMemoryContextPlanAware() string {
|
||||||
|
var parts []string
|
||||||
|
state := ms.getParsedPlanState()
|
||||||
|
longTerm := state.content
|
||||||
|
|
||||||
|
if longTerm != "" {
|
||||||
|
if state.hasActivePlan {
|
||||||
|
switch state.status {
|
||||||
|
case "interviewing":
|
||||||
|
parts = append(parts, ms.getInterviewContextFrom(longTerm))
|
||||||
|
case "review":
|
||||||
|
parts = append(parts, ms.getReviewContextFrom(longTerm))
|
||||||
|
default:
|
||||||
|
parts = append(parts, ms.getPlanContextFrom(longTerm))
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
parts = append(parts, "## Long-term Memory\n\n"+longTerm)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Suppress daily notes when a plan is active to save context
|
||||||
|
if !state.hasActivePlan {
|
||||||
|
recentNotes := ms.GetRecentDailyNotes(3)
|
||||||
|
if recentNotes != "" {
|
||||||
|
parts = append(parts, "## Recent Daily Notes\n\n"+recentNotes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(parts) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return strings.Join(parts, "\n\n---\n\n")
|
||||||
|
}
|
||||||
|
|
@ -33,6 +33,7 @@ type OutboundMessage struct {
|
||||||
Channel string `json:"channel"`
|
Channel string `json:"channel"`
|
||||||
ChatID string `json:"chat_id"`
|
ChatID string `json:"chat_id"`
|
||||||
Content string `json:"content"`
|
Content string `json:"content"`
|
||||||
|
ReplyToMessageID string `json:"reply_to_message_id,omitempty"`
|
||||||
IsStatus bool `json:"is_status,omitempty"`
|
IsStatus bool `json:"is_status,omitempty"`
|
||||||
IsTaskStatus bool `json:"is_task_status,omitempty"`
|
IsTaskStatus bool `json:"is_task_status,omitempty"`
|
||||||
TaskID string `json:"task_id,omitempty"`
|
TaskID string `json:"task_id,omitempty"`
|
||||||
|
|
|
||||||
|
|
@ -26,12 +26,6 @@ type ReactionCapable interface {
|
||||||
ReactToMessage(ctx context.Context, chatID, messageID string) (undo func(), err error)
|
ReactToMessage(ctx context.Context, chatID, messageID string) (undo func(), err error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// MessageSenderWithID — channels that can send a message and return its platform-specific ID.
|
|
||||||
// Used by Manager to track status/task messages for later editing.
|
|
||||||
type MessageSenderWithID interface {
|
|
||||||
SendWithID(ctx context.Context, chatID string, content string) (messageID string, err error)
|
|
||||||
}
|
|
||||||
|
|
||||||
// PlaceholderCapable — channels that can send a placeholder message
|
// PlaceholderCapable — channels that can send a placeholder message
|
||||||
// (e.g. "Thinking... 💭") that will later be edited to the actual response.
|
// (e.g. "Thinking... 💭") that will later be edited to the actual response.
|
||||||
// The channel MUST also implement MessageEditor for the placeholder to be useful.
|
// The channel MUST also implement MessageEditor for the placeholder to be useful.
|
||||||
|
|
@ -41,13 +35,6 @@ type PlaceholderCapable interface {
|
||||||
SendPlaceholder(ctx context.Context, chatID string) (messageID string, err error)
|
SendPlaceholder(ctx context.Context, chatID string) (messageID string, err error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// DraftSender — channels that can send progressive draft messages.
|
|
||||||
// Used for streaming LLM output without the "edited" indicator.
|
|
||||||
// draftID must be non-zero and consistent across updates for the same draft.
|
|
||||||
type DraftSender interface {
|
|
||||||
SendDraft(ctx context.Context, chatID string, draftID int, content string) error
|
|
||||||
}
|
|
||||||
|
|
||||||
// PlaceholderRecorder is injected into channels by Manager.
|
// PlaceholderRecorder is injected into channels by Manager.
|
||||||
// Channels call these methods on inbound to register typing/placeholder state.
|
// Channels call these methods on inbound to register typing/placeholder state.
|
||||||
// Manager uses the registered state on outbound to stop typing and edit placeholders.
|
// Manager uses the registered state on outbound to stop typing and edit placeholders.
|
||||||
|
|
|
||||||
16
pkg/channels/interfaces_ext.go
Normal file
16
pkg/channels/interfaces_ext.go
Normal file
|
|
@ -0,0 +1,16 @@
|
||||||
|
package channels
|
||||||
|
|
||||||
|
import "context"
|
||||||
|
|
||||||
|
// MessageSenderWithID — channels that can send a message and return its platform-specific ID.
|
||||||
|
// Used by Manager to track status/task messages for later editing.
|
||||||
|
type MessageSenderWithID interface {
|
||||||
|
SendWithID(ctx context.Context, chatID string, content string) (messageID string, err error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DraftSender — channels that can send progressive draft messages.
|
||||||
|
// Used for streaming LLM output without the "edited" indicator.
|
||||||
|
// draftID must be non-zero and consistent across updates for the same draft.
|
||||||
|
type DraftSender interface {
|
||||||
|
SendDraft(ctx context.Context, chatID string, draftID int, content string) error
|
||||||
|
}
|
||||||
|
|
@ -23,6 +23,14 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
dedupTTL = 5 * time.Minute
|
||||||
|
dedupInterval = 60 * time.Second
|
||||||
|
dedupMaxSize = 10000 // hard cap on dedup map entries
|
||||||
|
typingResend = 8 * time.Second
|
||||||
|
typingSeconds = 10
|
||||||
|
)
|
||||||
|
|
||||||
type QQChannel struct {
|
type QQChannel struct {
|
||||||
*channels.BaseChannel
|
*channels.BaseChannel
|
||||||
config config.QQConfig
|
config config.QQConfig
|
||||||
|
|
@ -31,12 +39,28 @@ type QQChannel struct {
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
sessionManager botgo.SessionManager
|
sessionManager botgo.SessionManager
|
||||||
processedIDs map[string]bool
|
|
||||||
mu sync.RWMutex
|
// Chat routing: track whether a chatID is group or direct.
|
||||||
|
chatType sync.Map // chatID → "group" | "direct"
|
||||||
|
|
||||||
|
// Passive reply: store last inbound message ID per chat.
|
||||||
|
lastMsgID sync.Map // chatID → string
|
||||||
|
|
||||||
|
// msg_seq: per-chat atomic counter for multi-part replies.
|
||||||
|
msgSeqCounters sync.Map // chatID → *atomic.Uint64
|
||||||
|
|
||||||
|
// Time-based dedup replacing the unbounded map.
|
||||||
|
dedup map[string]time.Time
|
||||||
|
muDedup sync.Mutex
|
||||||
|
|
||||||
|
// done is closed on Stop to shut down the dedup janitor.
|
||||||
|
done chan struct{}
|
||||||
|
stopOnce sync.Once
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewQQChannel(cfg config.QQConfig, messageBus *bus.MessageBus) (*QQChannel, error) {
|
func NewQQChannel(cfg config.QQConfig, messageBus *bus.MessageBus) (*QQChannel, error) {
|
||||||
base := channels.NewBaseChannel("qq", cfg, messageBus, cfg.AllowFrom,
|
base := channels.NewBaseChannel("qq", cfg, messageBus, cfg.AllowFrom,
|
||||||
|
channels.WithMaxMessageLength(cfg.MaxMessageLength),
|
||||||
channels.WithGroupTrigger(cfg.GroupTrigger),
|
channels.WithGroupTrigger(cfg.GroupTrigger),
|
||||||
channels.WithReasoningChannelID(cfg.ReasoningChannelID),
|
channels.WithReasoningChannelID(cfg.ReasoningChannelID),
|
||||||
)
|
)
|
||||||
|
|
@ -44,7 +68,8 @@ func NewQQChannel(cfg config.QQConfig, messageBus *bus.MessageBus) (*QQChannel,
|
||||||
return &QQChannel{
|
return &QQChannel{
|
||||||
BaseChannel: base,
|
BaseChannel: base,
|
||||||
config: cfg,
|
config: cfg,
|
||||||
processedIDs: make(map[string]bool),
|
dedup: make(map[string]time.Time),
|
||||||
|
done: make(chan struct{}),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -53,8 +78,13 @@ func (c *QQChannel) Start(ctx context.Context) error {
|
||||||
return fmt.Errorf("QQ app_id and app_secret not configured")
|
return fmt.Errorf("QQ app_id and app_secret not configured")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
botgo.SetLogger(logger.NewLogger("botgo"))
|
||||||
logger.InfoC("qq", "Starting QQ bot (WebSocket mode)")
|
logger.InfoC("qq", "Starting QQ bot (WebSocket mode)")
|
||||||
|
|
||||||
|
// Reinitialize shutdown signal for clean restart.
|
||||||
|
c.done = make(chan struct{})
|
||||||
|
c.stopOnce = sync.Once{}
|
||||||
|
|
||||||
// create token source
|
// create token source
|
||||||
credentials := &token.QQBotCredentials{
|
credentials := &token.QQBotCredentials{
|
||||||
AppID: c.config.AppID,
|
AppID: c.config.AppID,
|
||||||
|
|
@ -102,6 +132,15 @@ func (c *QQChannel) Start(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
|
// start dedup janitor goroutine
|
||||||
|
go c.dedupJanitor()
|
||||||
|
|
||||||
|
// Pre-register reasoning_channel_id as group chat if configured,
|
||||||
|
// so outbound-only destinations are routed correctly.
|
||||||
|
if c.config.ReasoningChannelID != "" {
|
||||||
|
c.chatType.Store(c.config.ReasoningChannelID, "group")
|
||||||
|
}
|
||||||
|
|
||||||
c.SetRunning(true)
|
c.SetRunning(true)
|
||||||
logger.InfoC("qq", "QQ bot started successfully")
|
logger.InfoC("qq", "QQ bot started successfully")
|
||||||
|
|
||||||
|
|
@ -112,6 +151,9 @@ func (c *QQChannel) Stop(ctx context.Context) error {
|
||||||
logger.InfoC("qq", "Stopping QQ bot")
|
logger.InfoC("qq", "Stopping QQ bot")
|
||||||
c.SetRunning(false)
|
c.SetRunning(false)
|
||||||
|
|
||||||
|
// Signal the dedup janitor to stop (idempotent).
|
||||||
|
c.stopOnce.Do(func() { close(c.done) })
|
||||||
|
|
||||||
if c.cancel != nil {
|
if c.cancel != nil {
|
||||||
c.cancel()
|
c.cancel()
|
||||||
}
|
}
|
||||||
|
|
@ -119,20 +161,81 @@ func (c *QQChannel) Stop(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// getChatKind returns the chat type for a given chatID ("group" or "direct").
|
||||||
|
// Unknown chatIDs default to "group" and log a warning, since QQ group IDs are
|
||||||
|
// more common as outbound-only destinations (e.g. reasoning_channel_id).
|
||||||
|
func (c *QQChannel) getChatKind(chatID string) string {
|
||||||
|
if v, ok := c.chatType.Load(chatID); ok {
|
||||||
|
if k, ok := v.(string); ok {
|
||||||
|
return k
|
||||||
|
}
|
||||||
|
}
|
||||||
|
logger.DebugCF("qq", "Unknown chat type for chatID, defaulting to group", map[string]any{
|
||||||
|
"chat_id": chatID,
|
||||||
|
})
|
||||||
|
return "group"
|
||||||
|
}
|
||||||
|
|
||||||
func (c *QQChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
func (c *QQChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
// construct message
|
chatKind := c.getChatKind(msg.ChatID)
|
||||||
|
|
||||||
|
// Build message with content.
|
||||||
msgToCreate := &dto.MessageToCreate{
|
msgToCreate := &dto.MessageToCreate{
|
||||||
Content: msg.Content,
|
Content: msg.Content,
|
||||||
|
MsgType: dto.TextMsg,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use Markdown message type if enabled in config.
|
||||||
|
if c.config.SendMarkdown {
|
||||||
|
msgToCreate.MsgType = dto.MarkdownMsg
|
||||||
|
msgToCreate.Markdown = &dto.Markdown{
|
||||||
|
Content: msg.Content,
|
||||||
|
}
|
||||||
|
// Clear plain content to avoid sending duplicate text.
|
||||||
|
msgToCreate.Content = ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Attach passive reply msg_id and msg_seq if available.
|
||||||
|
if v, ok := c.lastMsgID.Load(msg.ChatID); ok {
|
||||||
|
if msgID, ok := v.(string); ok && msgID != "" {
|
||||||
|
msgToCreate.MsgID = msgID
|
||||||
|
|
||||||
|
// Increment msg_seq atomically for multi-part replies.
|
||||||
|
if counterVal, ok := c.msgSeqCounters.Load(msg.ChatID); ok {
|
||||||
|
if counter, ok := counterVal.(*atomic.Uint64); ok {
|
||||||
|
seq := counter.Add(1)
|
||||||
|
msgToCreate.MsgSeq = uint32(seq)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sanitize URLs in group messages to avoid QQ's URL blacklist rejection.
|
||||||
|
if chatKind == "group" {
|
||||||
|
if msgToCreate.Content != "" {
|
||||||
|
msgToCreate.Content = sanitizeURLs(msgToCreate.Content)
|
||||||
|
}
|
||||||
|
if msgToCreate.Markdown != nil && msgToCreate.Markdown.Content != "" {
|
||||||
|
msgToCreate.Markdown.Content = sanitizeURLs(msgToCreate.Markdown.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Route to group or C2C.
|
||||||
|
var err error
|
||||||
|
if chatKind == "group" {
|
||||||
|
_, err = c.api.PostGroupMessage(ctx, msg.ChatID, msgToCreate)
|
||||||
|
} else {
|
||||||
|
_, err = c.api.PostC2CMessage(ctx, msg.ChatID, msgToCreate)
|
||||||
}
|
}
|
||||||
|
|
||||||
// send C2C message
|
|
||||||
_, err := c.api.PostC2CMessage(ctx, msg.ChatID, msgToCreate)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorCF("qq", "Failed to send C2C message", map[string]any{
|
logger.ErrorCF("qq", "Failed to send message", map[string]any{
|
||||||
|
"chat_id": msg.ChatID,
|
||||||
|
"chat_kind": chatKind,
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return fmt.Errorf("qq send: %w", channels.ErrTemporary)
|
return fmt.Errorf("qq send: %w", channels.ErrTemporary)
|
||||||
|
|
@ -141,7 +244,150 @@ func (c *QQChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleC2CMessage handles QQ private messages
|
// StartTyping implements channels.TypingCapable.
|
||||||
|
// It sends an InputNotify (msg_type=6) immediately and re-sends every 8 seconds.
|
||||||
|
// The returned stop function is idempotent and cancels the goroutine.
|
||||||
|
func (c *QQChannel) StartTyping(ctx context.Context, chatID string) (func(), error) {
|
||||||
|
// We need a stored msg_id for passive InputNotify; skip if none available.
|
||||||
|
v, ok := c.lastMsgID.Load(chatID)
|
||||||
|
if !ok {
|
||||||
|
return func() {}, nil
|
||||||
|
}
|
||||||
|
msgID, ok := v.(string)
|
||||||
|
if !ok || msgID == "" {
|
||||||
|
return func() {}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
chatKind := c.getChatKind(chatID)
|
||||||
|
|
||||||
|
sendTyping := func(sendCtx context.Context) {
|
||||||
|
typingMsg := &dto.MessageToCreate{
|
||||||
|
MsgType: dto.InputNotifyMsg,
|
||||||
|
MsgID: msgID,
|
||||||
|
InputNotify: &dto.InputNotify{
|
||||||
|
InputType: 1,
|
||||||
|
InputSecond: typingSeconds,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
var err error
|
||||||
|
if chatKind == "group" {
|
||||||
|
_, err = c.api.PostGroupMessage(sendCtx, chatID, typingMsg)
|
||||||
|
} else {
|
||||||
|
_, err = c.api.PostC2CMessage(sendCtx, chatID, typingMsg)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
logger.DebugCF("qq", "Failed to send typing indicator", map[string]any{
|
||||||
|
"chat_id": chatID,
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Send immediately.
|
||||||
|
sendTyping(c.ctx)
|
||||||
|
|
||||||
|
typingCtx, cancel := context.WithCancel(c.ctx)
|
||||||
|
go func() {
|
||||||
|
ticker := time.NewTicker(typingResend)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-typingCtx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
sendTyping(typingCtx)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return cancel, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendMedia implements the channels.MediaSender interface.
|
||||||
|
// QQ RichMediaMessage requires an HTTP/HTTPS URL — local file paths are not supported.
|
||||||
|
// If part.Ref is already an http(s) URL it is used directly; otherwise we try
|
||||||
|
// the media store, and skip with a warning if the resolved path is not an HTTP URL.
|
||||||
|
func (c *QQChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
|
||||||
|
if !c.IsRunning() {
|
||||||
|
return channels.ErrNotRunning
|
||||||
|
}
|
||||||
|
|
||||||
|
chatKind := c.getChatKind(msg.ChatID)
|
||||||
|
|
||||||
|
for _, part := range msg.Parts {
|
||||||
|
// If the ref is already an HTTP(S) URL, use it directly.
|
||||||
|
mediaURL := part.Ref
|
||||||
|
if !isHTTPURL(mediaURL) {
|
||||||
|
// Try resolving through media store.
|
||||||
|
store := c.GetMediaStore()
|
||||||
|
if store == nil {
|
||||||
|
logger.WarnCF("qq", "QQ media requires HTTP/HTTPS URL, no media store available", map[string]any{
|
||||||
|
"ref": part.Ref,
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
resolved, err := store.Resolve(part.Ref)
|
||||||
|
if err != nil {
|
||||||
|
logger.ErrorCF("qq", "Failed to resolve media ref", map[string]any{
|
||||||
|
"ref": part.Ref,
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if !isHTTPURL(resolved) {
|
||||||
|
logger.WarnCF("qq", "QQ media requires HTTP/HTTPS URL, local files not supported", map[string]any{
|
||||||
|
"ref": part.Ref,
|
||||||
|
"resolved": resolved,
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
mediaURL = resolved
|
||||||
|
}
|
||||||
|
|
||||||
|
// Map part type to QQ file type: 1=image, 2=video, 3=audio, 4=file.
|
||||||
|
var fileType uint64
|
||||||
|
switch part.Type {
|
||||||
|
case "image":
|
||||||
|
fileType = 1
|
||||||
|
case "video":
|
||||||
|
fileType = 2
|
||||||
|
case "audio":
|
||||||
|
fileType = 3
|
||||||
|
default:
|
||||||
|
fileType = 4 // file
|
||||||
|
}
|
||||||
|
|
||||||
|
richMedia := &dto.RichMediaMessage{
|
||||||
|
FileType: fileType,
|
||||||
|
URL: mediaURL,
|
||||||
|
SrvSendMsg: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
var sendErr error
|
||||||
|
if chatKind == "group" {
|
||||||
|
_, sendErr = c.api.PostGroupMessage(ctx, msg.ChatID, richMedia)
|
||||||
|
} else {
|
||||||
|
_, sendErr = c.api.PostC2CMessage(ctx, msg.ChatID, richMedia)
|
||||||
|
}
|
||||||
|
|
||||||
|
if sendErr != nil {
|
||||||
|
logger.ErrorCF("qq", "Failed to send media", map[string]any{
|
||||||
|
"type": part.Type,
|
||||||
|
"chat_id": msg.ChatID,
|
||||||
|
"error": sendErr.Error(),
|
||||||
|
})
|
||||||
|
return fmt.Errorf("qq send media: %w", channels.ErrTemporary)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleC2CMessage handles QQ private messages.
|
||||||
func (c *QQChannel) handleC2CMessage() event.C2CMessageEventHandler {
|
func (c *QQChannel) handleC2CMessage() event.C2CMessageEventHandler {
|
||||||
return func(event *dto.WSPayload, data *dto.WSC2CMessageData) error {
|
return func(event *dto.WSPayload, data *dto.WSC2CMessageData) error {
|
||||||
// deduplication check
|
// deduplication check
|
||||||
|
|
@ -170,7 +416,13 @@ func (c *QQChannel) handleC2CMessage() event.C2CMessageEventHandler {
|
||||||
"length": len(content),
|
"length": len(content),
|
||||||
})
|
})
|
||||||
|
|
||||||
// 转发到消息总线
|
// Store chat routing context.
|
||||||
|
c.chatType.Store(senderID, "direct")
|
||||||
|
c.lastMsgID.Store(senderID, data.ID)
|
||||||
|
|
||||||
|
// Reset msg_seq counter for new inbound message.
|
||||||
|
c.msgSeqCounters.Store(senderID, new(atomic.Uint64))
|
||||||
|
|
||||||
metadata := map[string]string{}
|
metadata := map[string]string{}
|
||||||
|
|
||||||
sender := bus.SenderInfo{
|
sender := bus.SenderInfo{
|
||||||
|
|
@ -198,7 +450,7 @@ func (c *QQChannel) handleC2CMessage() event.C2CMessageEventHandler {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleGroupATMessage handles QQ group @ messages
|
// handleGroupATMessage handles QQ group @ messages.
|
||||||
func (c *QQChannel) handleGroupATMessage() event.GroupATMessageEventHandler {
|
func (c *QQChannel) handleGroupATMessage() event.GroupATMessageEventHandler {
|
||||||
return func(event *dto.WSPayload, data *dto.WSGroupATMessageData) error {
|
return func(event *dto.WSPayload, data *dto.WSGroupATMessageData) error {
|
||||||
// deduplication check
|
// deduplication check
|
||||||
|
|
@ -235,7 +487,13 @@ func (c *QQChannel) handleGroupATMessage() event.GroupATMessageEventHandler {
|
||||||
"length": len(content),
|
"length": len(content),
|
||||||
})
|
})
|
||||||
|
|
||||||
// 转发到消息总线(使用 GroupID 作为 ChatID)
|
// Store chat routing context using GroupID as chatID.
|
||||||
|
c.chatType.Store(data.GroupID, "group")
|
||||||
|
c.lastMsgID.Store(data.GroupID, data.ID)
|
||||||
|
|
||||||
|
// Reset msg_seq counter for new inbound message.
|
||||||
|
c.msgSeqCounters.Store(data.GroupID, new(atomic.Uint64))
|
||||||
|
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"group_id": data.GroupID,
|
"group_id": data.GroupID,
|
||||||
}
|
}
|
||||||
|
|
@ -265,29 +523,102 @@ func (c *QQChannel) handleGroupATMessage() event.GroupATMessageEventHandler {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// isDuplicate 检查消息是否重复
|
// isDuplicate checks whether a message has been seen within the TTL window.
|
||||||
|
// It also enforces a hard cap on map size by evicting oldest entries.
|
||||||
func (c *QQChannel) isDuplicate(messageID string) bool {
|
func (c *QQChannel) isDuplicate(messageID string) bool {
|
||||||
c.mu.Lock()
|
c.muDedup.Lock()
|
||||||
defer c.mu.Unlock()
|
defer c.muDedup.Unlock()
|
||||||
|
|
||||||
if c.processedIDs[messageID] {
|
if ts, exists := c.dedup[messageID]; exists && time.Since(ts) < dedupTTL {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
c.processedIDs[messageID] = true
|
// Enforce hard cap: evict oldest entries when at capacity.
|
||||||
|
if len(c.dedup) >= dedupMaxSize {
|
||||||
// 简单清理:限制 map 大小
|
var oldestID string
|
||||||
if len(c.processedIDs) > 10000 {
|
var oldestTS time.Time
|
||||||
// 清空一半
|
for id, ts := range c.dedup {
|
||||||
count := 0
|
if oldestID == "" || ts.Before(oldestTS) {
|
||||||
for id := range c.processedIDs {
|
oldestID = id
|
||||||
if count >= 5000 {
|
oldestTS = ts
|
||||||
break
|
|
||||||
}
|
}
|
||||||
delete(c.processedIDs, id)
|
}
|
||||||
count++
|
if oldestID != "" {
|
||||||
|
delete(c.dedup, oldestID)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
c.dedup[messageID] = time.Now()
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// dedupJanitor periodically evicts expired entries from the dedup map.
|
||||||
|
func (c *QQChannel) dedupJanitor() {
|
||||||
|
ticker := time.NewTicker(dedupInterval)
|
||||||
|
defer ticker.Stop()
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-c.done:
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
// Collect expired keys under read-like scan.
|
||||||
|
c.muDedup.Lock()
|
||||||
|
now := time.Now()
|
||||||
|
var expired []string
|
||||||
|
for id, ts := range c.dedup {
|
||||||
|
if now.Sub(ts) >= dedupTTL {
|
||||||
|
expired = append(expired, id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, id := range expired {
|
||||||
|
delete(c.dedup, id)
|
||||||
|
}
|
||||||
|
c.muDedup.Unlock()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// isHTTPURL returns true if s starts with http:// or https://.
|
||||||
|
func isHTTPURL(s string) bool {
|
||||||
|
return strings.HasPrefix(s, "http://") || strings.HasPrefix(s, "https://")
|
||||||
|
}
|
||||||
|
|
||||||
|
// urlPattern matches URLs with explicit http(s):// scheme.
|
||||||
|
// Only scheme-prefixed URLs are matched to avoid false positives on bare text
|
||||||
|
// like version numbers (e.g., "1.2.3") or domain-like fragments.
|
||||||
|
var urlPattern = regexp.MustCompile(
|
||||||
|
`(?i)` +
|
||||||
|
`https?://` + // required scheme
|
||||||
|
`(?:[a-zA-Z0-9](?:[a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?\.)+` + // domain parts
|
||||||
|
`[a-zA-Z]{2,}` + // TLD
|
||||||
|
`(?:[/?#]\S*)?`, // optional path/query/fragment
|
||||||
|
)
|
||||||
|
|
||||||
|
// sanitizeURLs replaces dots in URL domains with "。" (fullwidth period)
|
||||||
|
// to prevent QQ's URL blacklist from rejecting the message.
|
||||||
|
func sanitizeURLs(text string) string {
|
||||||
|
return urlPattern.ReplaceAllStringFunc(text, func(match string) string {
|
||||||
|
// Split into scheme + rest (scheme is always present).
|
||||||
|
idx := strings.Index(match, "://")
|
||||||
|
scheme := match[:idx+3]
|
||||||
|
rest := match[idx+3:]
|
||||||
|
|
||||||
|
// Find where the domain ends (first / ? or #).
|
||||||
|
domainEnd := len(rest)
|
||||||
|
for i, ch := range rest {
|
||||||
|
if ch == '/' || ch == '?' || ch == '#' {
|
||||||
|
domainEnd = i
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
domain := rest[:domainEnd]
|
||||||
|
path := rest[domainEnd:]
|
||||||
|
|
||||||
|
// Replace dots in domain only.
|
||||||
|
domain = strings.ReplaceAll(domain, ".", "。")
|
||||||
|
|
||||||
|
return scheme + domain + path
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -77,6 +77,7 @@ func NewTelegramChannel(cfg *config.Config, bus *bus.MessageBus) (*TelegramChann
|
||||||
if baseURL := strings.TrimRight(strings.TrimSpace(telegramCfg.BaseURL), "/"); baseURL != "" {
|
if baseURL := strings.TrimRight(strings.TrimSpace(telegramCfg.BaseURL), "/"); baseURL != "" {
|
||||||
opts = append(opts, telego.WithAPIServer(baseURL))
|
opts = append(opts, telego.WithAPIServer(baseURL))
|
||||||
}
|
}
|
||||||
|
opts = append(opts, telego.WithLogger(logger.NewLogger("telego")))
|
||||||
|
|
||||||
bot, err := telego.NewBot(telegramCfg.Token, opts...)
|
bot, err := telego.NewBot(telegramCfg.Token, opts...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -168,7 +169,7 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
return channels.ErrNotRunning
|
return channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
chatID, threadID, err := parseChatID(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 fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
@ -180,6 +181,7 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
// 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
|
||||||
queue := []string{msg.Content}
|
queue := []string{msg.Content}
|
||||||
for len(queue) > 0 {
|
for len(queue) > 0 {
|
||||||
chunk := queue[0]
|
chunk := queue[0]
|
||||||
|
|
@ -200,9 +202,11 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.sendHTMLChunk(ctx, chatID, htmlContent, chunk); err != nil {
|
if err := c.sendHTMLChunk(ctx, chatID, threadID, htmlContent, chunk, replyToID); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
// Only the first chunk should be a reply; subsequent chunks are normal messages.
|
||||||
|
replyToID = ""
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
|
|
@ -210,11 +214,19 @@ func (c *TelegramChannel) Send(ctx context.Context, msg bus.OutboundMessage) err
|
||||||
|
|
||||||
// sendHTMLChunk sends a single HTML message, falling back to the original
|
// sendHTMLChunk sends a single HTML message, falling back to the original
|
||||||
// markdown as plain text on parse failure so users never see raw HTML tags.
|
// markdown as plain text on parse failure so users never see raw HTML tags.
|
||||||
func (c *TelegramChannel) sendHTMLChunk(ctx context.Context, chatID int64, htmlContent, mdFallback string) error {
|
func (c *TelegramChannel) sendHTMLChunk(
|
||||||
|
ctx context.Context, chatID int64, threadID int, htmlContent, mdFallback string, replyToID string,
|
||||||
|
) error {
|
||||||
tgMsg := tu.Message(tu.ID(chatID), htmlContent)
|
tgMsg := tu.Message(tu.ID(chatID), htmlContent)
|
||||||
tgMsg.ParseMode = telego.ModeHTML
|
tgMsg.ParseMode = telego.ModeHTML
|
||||||
if threadID != 0 {
|
|
||||||
tgMsg.MessageThreadID = threadID
|
tgMsg.MessageThreadID = threadID
|
||||||
|
|
||||||
|
if replyToID != "" {
|
||||||
|
if mid, parseErr := strconv.Atoi(replyToID); parseErr == nil {
|
||||||
|
tgMsg.ReplyParameters = &telego.ReplyParameters{
|
||||||
|
MessageID: mid,
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := c.bot.SendMessage(ctx, tgMsg); err != nil {
|
if _, err := c.bot.SendMessage(ctx, tgMsg); err != nil {
|
||||||
|
|
@ -230,54 +242,21 @@ func (c *TelegramChannel) sendHTMLChunk(ctx context.Context, chatID int64, htmlC
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendWithID implements channels.MessageSenderWithID.
|
|
||||||
// It sends a message and returns the platform message ID.
|
|
||||||
func (c *TelegramChannel) SendWithID(ctx context.Context, chatID string, content string) (string, error) {
|
|
||||||
if !c.IsRunning() {
|
|
||||||
return "", channels.ErrNotRunning
|
|
||||||
}
|
|
||||||
|
|
||||||
cid, tid, err := parseChatID(chatID)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("invalid chat ID %s: %w", chatID, channels.ErrSendFailed)
|
|
||||||
}
|
|
||||||
|
|
||||||
htmlContent := markdownToTelegramHTML(content)
|
|
||||||
tgMsg := tu.Message(tu.ID(cid), htmlContent)
|
|
||||||
tgMsg.ParseMode = telego.ModeHTML
|
|
||||||
if tid != 0 {
|
|
||||||
tgMsg.MessageThreadID = tid
|
|
||||||
}
|
|
||||||
|
|
||||||
sent, err := c.bot.SendMessage(ctx, tgMsg)
|
|
||||||
if err != nil {
|
|
||||||
// Fallback to plain text
|
|
||||||
tgMsg.ParseMode = ""
|
|
||||||
sent, err = c.bot.SendMessage(ctx, tgMsg)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("telegram send: %w", channels.ErrTemporary)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return fmt.Sprintf("%d", sent.MessageID), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// StartTyping implements channels.TypingCapable.
|
// StartTyping implements channels.TypingCapable.
|
||||||
// It sends ChatAction(typing) immediately and then repeats every 4 seconds
|
// It sends ChatAction(typing) immediately and then repeats every 4 seconds
|
||||||
// (Telegram's typing indicator expires after ~5s) in a background goroutine.
|
// (Telegram's typing indicator expires after ~5s) in a background goroutine.
|
||||||
// The returned stop function is idempotent and cancels the goroutine.
|
// The returned stop function is idempotent and cancels the goroutine.
|
||||||
func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(), error) {
|
func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(), error) {
|
||||||
cid, tid, err := parseChatID(chatID)
|
cid, threadID, err := parseTelegramChatID(chatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return func() {}, err
|
return func() {}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
action := tu.ChatAction(tu.ID(cid), telego.ChatActionTyping)
|
||||||
|
action.MessageThreadID = threadID
|
||||||
|
|
||||||
// Send the first typing action immediately
|
// Send the first typing action immediately
|
||||||
firstAction := tu.ChatAction(tu.ID(cid), telego.ChatActionTyping)
|
_ = c.bot.SendChatAction(ctx, action)
|
||||||
if tid != 0 {
|
|
||||||
firstAction.MessageThreadID = tid
|
|
||||||
}
|
|
||||||
_ = c.bot.SendChatAction(ctx, firstAction)
|
|
||||||
|
|
||||||
typingCtx, cancel := context.WithCancel(ctx)
|
typingCtx, cancel := context.WithCancel(ctx)
|
||||||
go func() {
|
go func() {
|
||||||
|
|
@ -288,11 +267,9 @@ func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(
|
||||||
case <-typingCtx.Done():
|
case <-typingCtx.Done():
|
||||||
return
|
return
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
action := tu.ChatAction(tu.ID(cid), telego.ChatActionTyping)
|
a := tu.ChatAction(tu.ID(cid), telego.ChatActionTyping)
|
||||||
if tid != 0 {
|
a.MessageThreadID = threadID
|
||||||
action.MessageThreadID = tid
|
_ = c.bot.SendChatAction(typingCtx, a)
|
||||||
}
|
|
||||||
_ = c.bot.SendChatAction(typingCtx, action)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
@ -302,7 +279,7 @@ func (c *TelegramChannel) StartTyping(ctx context.Context, chatID string) (func(
|
||||||
|
|
||||||
// EditMessage implements channels.MessageEditor.
|
// EditMessage implements channels.MessageEditor.
|
||||||
func (c *TelegramChannel) EditMessage(ctx context.Context, chatID string, messageID string, content string) error {
|
func (c *TelegramChannel) EditMessage(ctx context.Context, chatID string, messageID string, content string) error {
|
||||||
cid, _, err := parseChatID(chatID)
|
cid, _, err := parseTelegramChatID(chatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
@ -331,16 +308,14 @@ func (c *TelegramChannel) SendPlaceholder(ctx context.Context, chatID string) (s
|
||||||
text = "Thinking... 💭"
|
text = "Thinking... 💭"
|
||||||
}
|
}
|
||||||
|
|
||||||
cid, tid, err := parseChatID(chatID)
|
cid, threadID, err := parseTelegramChatID(chatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
params := tu.Message(tu.ID(cid), text)
|
phMsg := tu.Message(tu.ID(cid), text)
|
||||||
if tid != 0 {
|
phMsg.MessageThreadID = threadID
|
||||||
params.MessageThreadID = tid
|
pMsg, err := c.bot.SendMessage(ctx, phMsg)
|
||||||
}
|
|
||||||
pMsg, err := c.bot.SendMessage(ctx, params)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
@ -348,44 +323,13 @@ func (c *TelegramChannel) SendPlaceholder(ctx context.Context, chatID string) (s
|
||||||
return fmt.Sprintf("%d", pMsg.MessageID), nil
|
return fmt.Sprintf("%d", pMsg.MessageID), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendDraft implements channels.DraftSender.
|
|
||||||
// It uses Telegram Bot API's sendMessageDraft for progressive message streaming
|
|
||||||
// without the "edited" indicator. In groups, draft is used for dedicated topics only.
|
|
||||||
func (c *TelegramChannel) SendDraft(ctx context.Context, chatID string, draftID int, content string) error {
|
|
||||||
if !c.IsRunning() {
|
|
||||||
return channels.ErrNotRunning
|
|
||||||
}
|
|
||||||
cid, tid, err := parseChatID(chatID)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("invalid chat ID %s: %w", chatID, channels.ErrSendFailed)
|
|
||||||
}
|
|
||||||
if !isLikelyPrivateChatID(cid) && tid == 0 {
|
|
||||||
return fmt.Errorf("telegram draft unsupported for non-threaded group chat: %w", channels.ErrSendFailed)
|
|
||||||
}
|
|
||||||
htmlContent := markdownToTelegramHTML(content)
|
|
||||||
params := &telego.SendMessageDraftParams{
|
|
||||||
ChatID: cid,
|
|
||||||
MessageThreadID: tid,
|
|
||||||
DraftID: draftID,
|
|
||||||
Text: htmlContent,
|
|
||||||
ParseMode: telego.ModeHTML,
|
|
||||||
}
|
|
||||||
if err = c.bot.SendMessageDraft(ctx, params); err != nil {
|
|
||||||
// HTML parse failure — retry as plain text
|
|
||||||
params.ParseMode = ""
|
|
||||||
params.Text = content
|
|
||||||
return c.bot.SendMessageDraft(ctx, params)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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) error {
|
||||||
if !c.IsRunning() {
|
if !c.IsRunning() {
|
||||||
return channels.ErrNotRunning
|
return channels.ErrNotRunning
|
||||||
}
|
}
|
||||||
|
|
||||||
chatID, threadID, err := parseChatID(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 fmt.Errorf("invalid chat ID %s: %w", msg.ChatID, channels.ErrSendFailed)
|
||||||
}
|
}
|
||||||
|
|
@ -492,12 +436,11 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
|
|
||||||
chatID := message.Chat.ID
|
chatID := message.Chat.ID
|
||||||
c.chatIDs[platformID] = chatID
|
c.chatIDs[platformID] = chatID
|
||||||
threadID := message.MessageThreadID
|
|
||||||
|
|
||||||
content := ""
|
content := ""
|
||||||
mediaPaths := []string{}
|
mediaPaths := []string{}
|
||||||
|
|
||||||
chatIDStr := formatChatID(chatID, threadID)
|
chatIDStr := fmt.Sprintf("%d", chatID)
|
||||||
messageIDStr := fmt.Sprintf("%d", message.MessageID)
|
messageIDStr := fmt.Sprintf("%d", message.MessageID)
|
||||||
scope := channels.BuildMediaScope("telegram", chatIDStr, messageIDStr)
|
scope := channels.BuildMediaScope("telegram", chatIDStr, messageIDStr)
|
||||||
|
|
||||||
|
|
@ -589,21 +532,28 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
content = cleaned
|
content = cleaned
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// For forum topics, embed the thread ID as "chatID/threadID" so replies
|
||||||
|
// route to the correct topic and each topic gets its own session.
|
||||||
|
// Only forum groups (IsForum) are handled; regular group reply threads
|
||||||
|
// must share one session per group.
|
||||||
|
compositeChatID := fmt.Sprintf("%d", chatID)
|
||||||
|
threadID := message.MessageThreadID
|
||||||
|
if message.Chat.IsForum && threadID != 0 {
|
||||||
|
compositeChatID = fmt.Sprintf("%d/%d", chatID, threadID)
|
||||||
|
}
|
||||||
|
|
||||||
logger.DebugCF("telegram", "Received message", map[string]any{
|
logger.DebugCF("telegram", "Received message", map[string]any{
|
||||||
"sender_id": sender.CanonicalID,
|
"sender_id": sender.CanonicalID,
|
||||||
"chat_id": fmt.Sprintf("%d", chatID),
|
"chat_id": compositeChatID,
|
||||||
"thread_id": threadID,
|
"thread_id": threadID,
|
||||||
"chat_route": chatIDStr,
|
|
||||||
"preview": utils.Truncate(content, 50),
|
"preview": utils.Truncate(content, 50),
|
||||||
})
|
})
|
||||||
|
|
||||||
// Placeholder is now auto-triggered by BaseChannel.HandleMessage via PlaceholderCapable
|
|
||||||
|
|
||||||
peerKind := "direct"
|
peerKind := "direct"
|
||||||
peerID := fmt.Sprintf("%d", user.ID)
|
peerID := fmt.Sprintf("%d", user.ID)
|
||||||
if message.Chat.Type != "private" {
|
if message.Chat.Type != "private" {
|
||||||
peerKind = "group"
|
peerKind = "group"
|
||||||
peerID = fmt.Sprintf("%d", chatID)
|
peerID = compositeChatID
|
||||||
}
|
}
|
||||||
|
|
||||||
peer := bus.Peer{Kind: peerKind, ID: peerID}
|
peer := bus.Peer{Kind: peerKind, ID: peerID}
|
||||||
|
|
@ -616,11 +566,17 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Set parent_peer metadata for per-topic agent binding.
|
||||||
|
if message.Chat.IsForum && threadID != 0 {
|
||||||
|
metadata["parent_peer_kind"] = "topic"
|
||||||
|
metadata["parent_peer_id"] = fmt.Sprintf("%d", threadID)
|
||||||
|
}
|
||||||
|
|
||||||
c.HandleMessage(c.ctx,
|
c.HandleMessage(c.ctx,
|
||||||
peer,
|
peer,
|
||||||
messageID,
|
messageID,
|
||||||
platformID,
|
platformID,
|
||||||
chatIDStr,
|
compositeChatID,
|
||||||
content,
|
content,
|
||||||
mediaPaths,
|
mediaPaths,
|
||||||
metadata,
|
metadata,
|
||||||
|
|
@ -668,50 +624,25 @@ func (c *TelegramChannel) downloadFile(ctx context.Context, fileID, ext string)
|
||||||
return c.downloadFileWithInfo(file, ext)
|
return c.downloadFileWithInfo(file, ext)
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseChatID(chatIDStr string) (int64, int, error) {
|
// parseTelegramChatID splits "chatID/threadID" into its components.
|
||||||
trimmed := strings.TrimSpace(chatIDStr)
|
// Returns threadID=0 when no "/" is present (non-forum messages).
|
||||||
if trimmed == "" {
|
func parseTelegramChatID(chatID string) (int64, int, error) {
|
||||||
return 0, 0, fmt.Errorf("empty chat ID")
|
idx := strings.Index(chatID, "/")
|
||||||
|
if idx == -1 {
|
||||||
|
cid, err := strconv.ParseInt(chatID, 10, 64)
|
||||||
|
return cid, 0, err
|
||||||
}
|
}
|
||||||
|
cid, err := strconv.ParseInt(chatID[:idx], 10, 64)
|
||||||
parts := strings.Split(trimmed, "/")
|
|
||||||
if len(parts) > 2 {
|
|
||||||
return 0, 0, fmt.Errorf("invalid chat ID format: %q", chatIDStr)
|
|
||||||
}
|
|
||||||
|
|
||||||
cid, err := strconv.ParseInt(parts[0], 10, 64)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, 0, fmt.Errorf("invalid chat ID %q: %w", parts[0], err)
|
return 0, 0, err
|
||||||
}
|
}
|
||||||
|
tid, err := strconv.Atoi(chatID[idx+1:])
|
||||||
tid := 0
|
|
||||||
if len(parts) == 2 {
|
|
||||||
if parts[1] == "" {
|
|
||||||
return 0, 0, fmt.Errorf("invalid thread ID in %q", chatIDStr)
|
|
||||||
}
|
|
||||||
tid, err = strconv.Atoi(parts[1])
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, 0, fmt.Errorf("invalid thread ID %q: %w", parts[1], err)
|
return 0, 0, fmt.Errorf("invalid thread ID in chat ID %q: %w", chatID, err)
|
||||||
}
|
}
|
||||||
if tid < 0 {
|
|
||||||
return 0, 0, fmt.Errorf("thread ID must be non-negative: %d", tid)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return cid, tid, nil
|
return cid, tid, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func formatChatID(chatID int64, threadID int) string {
|
|
||||||
if threadID != 0 {
|
|
||||||
return fmt.Sprintf("%d/%d", chatID, threadID)
|
|
||||||
}
|
|
||||||
return fmt.Sprintf("%d", chatID)
|
|
||||||
}
|
|
||||||
|
|
||||||
func isLikelyPrivateChatID(chatID int64) bool {
|
|
||||||
return chatID > 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func markdownToTelegramHTML(text string) string {
|
func markdownToTelegramHTML(text string) string {
|
||||||
if text == "" {
|
if text == "" {
|
||||||
return ""
|
return ""
|
||||||
|
|
|
||||||
85
pkg/channels/telegram/telegram_ext.go
Normal file
85
pkg/channels/telegram/telegram_ext.go
Normal file
|
|
@ -0,0 +1,85 @@
|
||||||
|
package telegram
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/mymmrac/telego"
|
||||||
|
tu "github.com/mymmrac/telego/telegoutil"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/channels"
|
||||||
|
)
|
||||||
|
|
||||||
|
// SendWithID implements channels.MessageSenderWithID.
|
||||||
|
// It sends a message and returns the platform message ID.
|
||||||
|
func (c *TelegramChannel) SendWithID(ctx context.Context, chatID string, content string) (string, error) {
|
||||||
|
if !c.IsRunning() {
|
||||||
|
return "", channels.ErrNotRunning
|
||||||
|
}
|
||||||
|
|
||||||
|
cid, tid, err := parseTelegramChatID(chatID)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("invalid chat ID %s: %w", chatID, channels.ErrSendFailed)
|
||||||
|
}
|
||||||
|
|
||||||
|
htmlContent := markdownToTelegramHTML(content)
|
||||||
|
tgMsg := tu.Message(tu.ID(cid), htmlContent)
|
||||||
|
tgMsg.ParseMode = telego.ModeHTML
|
||||||
|
tgMsg.MessageThreadID = tid
|
||||||
|
|
||||||
|
sent, err := c.bot.SendMessage(ctx, tgMsg)
|
||||||
|
if err != nil {
|
||||||
|
// Fallback to plain text
|
||||||
|
tgMsg.ParseMode = ""
|
||||||
|
sent, err = c.bot.SendMessage(ctx, tgMsg)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("telegram send: %w", channels.ErrTemporary)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Sprintf("%d", sent.MessageID), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendDraft implements channels.DraftSender.
|
||||||
|
// It uses Telegram Bot API's sendMessageDraft for progressive message streaming
|
||||||
|
// without the "edited" indicator. In groups, draft is used for dedicated topics only.
|
||||||
|
func (c *TelegramChannel) SendDraft(ctx context.Context, chatID string, draftID int, content string) error {
|
||||||
|
if !c.IsRunning() {
|
||||||
|
return channels.ErrNotRunning
|
||||||
|
}
|
||||||
|
cid, tid, err := parseTelegramChatID(chatID)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid chat ID %s: %w", chatID, channels.ErrSendFailed)
|
||||||
|
}
|
||||||
|
if !isLikelyPrivateChatID(cid) && tid == 0 {
|
||||||
|
return fmt.Errorf("telegram draft unsupported for non-threaded group chat: %w", channels.ErrSendFailed)
|
||||||
|
}
|
||||||
|
htmlContent := markdownToTelegramHTML(content)
|
||||||
|
params := &telego.SendMessageDraftParams{
|
||||||
|
ChatID: cid,
|
||||||
|
MessageThreadID: tid,
|
||||||
|
DraftID: draftID,
|
||||||
|
Text: htmlContent,
|
||||||
|
ParseMode: telego.ModeHTML,
|
||||||
|
}
|
||||||
|
if err = c.bot.SendMessageDraft(ctx, params); err != nil {
|
||||||
|
// HTML parse failure — retry as plain text
|
||||||
|
params.ParseMode = ""
|
||||||
|
params.Text = content
|
||||||
|
return c.bot.SendMessageDraft(ctx, params)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// formatChatID formats a chat ID with optional thread ID as "chatID/threadID".
|
||||||
|
func formatChatID(chatID int64, threadID int) string {
|
||||||
|
if threadID != 0 {
|
||||||
|
return fmt.Sprintf("%d/%d", chatID, threadID)
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%d", chatID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// isLikelyPrivateChatID returns true for positive chat IDs (private chats).
|
||||||
|
func isLikelyPrivateChatID(chatID int64) bool {
|
||||||
|
return chatID > 0
|
||||||
|
}
|
||||||
|
|
@ -4,7 +4,7 @@ import (
|
||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestParseChatID(t *testing.T) {
|
func TestParseTelegramChatID(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
input string
|
input string
|
||||||
|
|
@ -14,29 +14,24 @@ func TestParseChatID(t *testing.T) {
|
||||||
}{
|
}{
|
||||||
{name: "plain private", input: "12345", wantCID: 12345, wantTID: 0},
|
{name: "plain private", input: "12345", wantCID: 12345, wantTID: 0},
|
||||||
{name: "group topic", input: "-100123/45", wantCID: -100123, wantTID: 45},
|
{name: "group topic", input: "-100123/45", wantCID: -100123, wantTID: 45},
|
||||||
{name: "trim spaces", input: " -100200/7 ", wantCID: -100200, wantTID: 7},
|
|
||||||
{name: "topic zero", input: "-100/0", wantCID: -100, wantTID: 0},
|
{name: "topic zero", input: "-100/0", wantCID: -100, wantTID: 0},
|
||||||
{name: "empty", input: "", wantErr: true},
|
|
||||||
{name: "bad chat", input: "abc/def", wantErr: true},
|
{name: "bad chat", input: "abc/def", wantErr: true},
|
||||||
{name: "missing topic", input: "-100/", wantErr: true},
|
|
||||||
{name: "too many parts", input: "-100/1/2", wantErr: true},
|
|
||||||
{name: "negative topic", input: "-100/-1", wantErr: true},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tc := range tests {
|
for _, tc := range tests {
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
gotCID, gotTID, err := parseChatID(tc.input)
|
gotCID, gotTID, err := parseTelegramChatID(tc.input)
|
||||||
if tc.wantErr {
|
if tc.wantErr {
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Fatalf("parseChatID(%q) expected error, got nil", tc.input)
|
t.Fatalf("parseTelegramChatID(%q) expected error, got nil", tc.input)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("parseChatID(%q) unexpected error: %v", tc.input, err)
|
t.Fatalf("parseTelegramChatID(%q) unexpected error: %v", tc.input, err)
|
||||||
}
|
}
|
||||||
if gotCID != tc.wantCID || gotTID != tc.wantTID {
|
if gotCID != tc.wantCID || gotTID != tc.wantTID {
|
||||||
t.Fatalf("parseChatID(%q) = (%d, %d), want (%d, %d)", tc.input, gotCID, gotTID, tc.wantCID, tc.wantTID)
|
t.Fatalf("parseTelegramChatID(%q) = (%d, %d), want (%d, %d)", tc.input, gotCID, gotTID, tc.wantCID, tc.wantTID)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -2,8 +2,6 @@ package providers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
|
||||||
"regexp"
|
|
||||||
"strings"
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -60,261 +58,6 @@ func extractToolCallsFromText(text string) []ToolCall {
|
||||||
return result
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// --- Shared helpers for XML tool call extraction ---
|
|
||||||
|
|
||||||
// normalizeAlpha keeps only lowercase ASCII letters.
|
|
||||||
// "tool_call" → "toolcall", "Tool-Call" → "toolcall", "ReadFile" → "readfile".
|
|
||||||
func normalizeAlpha(s string) string {
|
|
||||||
var b strings.Builder
|
|
||||||
for _, r := range s {
|
|
||||||
if r >= 'A' && r <= 'Z' {
|
|
||||||
b.WriteRune(r + 32)
|
|
||||||
} else if r >= 'a' && r <= 'z' {
|
|
||||||
b.WriteRune(r)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return b.String()
|
|
||||||
}
|
|
||||||
|
|
||||||
// levenshtein computes the edit distance between two strings.
|
|
||||||
// O(n*m) where n,m are string lengths — negligible for short tag names.
|
|
||||||
func levenshtein(a, b string) int {
|
|
||||||
la, lb := len(a), len(b)
|
|
||||||
if la == 0 {
|
|
||||||
return lb
|
|
||||||
}
|
|
||||||
if lb == 0 {
|
|
||||||
return la
|
|
||||||
}
|
|
||||||
prev := make([]int, lb+1)
|
|
||||||
for j := range prev {
|
|
||||||
prev[j] = j
|
|
||||||
}
|
|
||||||
for i := 1; i <= la; i++ {
|
|
||||||
curr := make([]int, lb+1)
|
|
||||||
curr[0] = i
|
|
||||||
for j := 1; j <= lb; j++ {
|
|
||||||
cost := 1
|
|
||||||
if a[i-1] == b[j-1] {
|
|
||||||
cost = 0
|
|
||||||
}
|
|
||||||
curr[j] = min(curr[j-1]+1, min(prev[j]+1, prev[j-1]+cost))
|
|
||||||
}
|
|
||||||
prev = curr
|
|
||||||
}
|
|
||||||
return prev[lb]
|
|
||||||
}
|
|
||||||
|
|
||||||
// Known tool call tag patterns (already alpha-normalized).
|
|
||||||
// Providers may use different names: tool_call, function_call, tool_use, etc.
|
|
||||||
var toolCallPatterns = []string{"toolcall", "functioncall", "tooluse"}
|
|
||||||
|
|
||||||
// isToolCallTag returns true if the tag name is close to any known tool call
|
|
||||||
// pattern after alpha normalization + edit distance (threshold ≤ 2).
|
|
||||||
func isToolCallTag(name string) bool {
|
|
||||||
const threshold = 2
|
|
||||||
norm := normalizeAlpha(name)
|
|
||||||
for _, pat := range toolCallPatterns {
|
|
||||||
if levenshtein(norm, pat) <= threshold {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
// tagSuffix returns the part after the last ':' (namespace separator),
|
|
||||||
// or the whole string if there is no ':'.
|
|
||||||
func tagSuffix(tag string) string {
|
|
||||||
if i := strings.LastIndex(tag, ":"); i >= 0 {
|
|
||||||
return tag[i+1:]
|
|
||||||
}
|
|
||||||
return tag
|
|
||||||
}
|
|
||||||
|
|
||||||
// --- XML block detection via regex ---
|
|
||||||
//
|
|
||||||
// Strategy: find <TAG>…</TAG> pairs using regex, then check if the tag
|
|
||||||
// suffix normalizes to something close to "toolcall" (edit distance ≤ 2).
|
|
||||||
// Uses greedy (longest) match for the closing tag to capture the full block.
|
|
||||||
|
|
||||||
var (
|
|
||||||
reOpenTag = regexp.MustCompile(`<([a-zA-Z][\w:.-]*)>`)
|
|
||||||
reCloseTag = regexp.MustCompile(`</([a-zA-Z][\w:.-]*)>`)
|
|
||||||
reBracketMarker = regexp.MustCompile(`\[TOOLCALL\]`)
|
|
||||||
)
|
|
||||||
|
|
||||||
// findToolCallBlock finds the first XML block whose tag suffix matches
|
|
||||||
// "toolcall" by edit distance. Returns the block boundaries and the inner
|
|
||||||
// content, or found=false.
|
|
||||||
func findToolCallBlock(text string) (blockStart, blockEnd int, content string, found bool) {
|
|
||||||
for _, om := range reOpenTag.FindAllStringSubmatchIndex(text, -1) {
|
|
||||||
tagName := text[om[2]:om[3]]
|
|
||||||
if !isToolCallTag(tagSuffix(tagName)) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
// Found a toolcall opening tag. Search for the last matching close tag (greedy).
|
|
||||||
afterOpen := text[om[1]:]
|
|
||||||
closes := reCloseTag.FindAllStringSubmatchIndex(afterOpen, -1)
|
|
||||||
for i := len(closes) - 1; i >= 0; i-- {
|
|
||||||
closeTagName := afterOpen[closes[i][2]:closes[i][3]]
|
|
||||||
if isToolCallTag(tagSuffix(closeTagName)) {
|
|
||||||
return om[0], om[1] + closes[i][1], afterOpen[:closes[i][0]], true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Fallback: look for orphaned closing tags (missing opening tag).
|
|
||||||
// Some LLMs emit the closing </ns:tool_call> without a matching opener.
|
|
||||||
// Reconstruct the block start from the first <invoke preceding the closer.
|
|
||||||
for _, cm := range reCloseTag.FindAllStringSubmatchIndex(text, -1) {
|
|
||||||
closeTagName := text[cm[2]:cm[3]]
|
|
||||||
if !isToolCallTag(tagSuffix(closeTagName)) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
// Found an orphaned toolcall closing tag. Scan backwards for <invoke.
|
|
||||||
before := text[:cm[0]]
|
|
||||||
invokePos := strings.LastIndex(before, "<invoke")
|
|
||||||
if invokePos == -1 {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
// Also consume a preceding [TOOLCALL] marker if present.
|
|
||||||
start := invokePos
|
|
||||||
if loc := reBracketMarker.FindStringIndex(before[:start]); loc != nil &&
|
|
||||||
strings.TrimSpace(before[loc[1]:start]) == "" {
|
|
||||||
start = loc[0]
|
|
||||||
}
|
|
||||||
return start, cm[1], text[invokePos:cm[0]], true
|
|
||||||
}
|
|
||||||
|
|
||||||
return 0, 0, "", false
|
|
||||||
}
|
|
||||||
|
|
||||||
// ExtractXMLToolCalls extracts tool calls from XML-formatted text.
|
|
||||||
//
|
|
||||||
// Expected format:
|
|
||||||
//
|
|
||||||
// <ns:toolcall>
|
|
||||||
// <invoke name="tool_name">
|
|
||||||
// <parameter name="param">value</parameter>
|
|
||||||
// </invoke>
|
|
||||||
// </ns:toolcall>
|
|
||||||
func ExtractXMLToolCalls(text string) []ToolCall {
|
|
||||||
return extractXMLToolCalls(text)
|
|
||||||
}
|
|
||||||
|
|
||||||
func extractXMLToolCalls(text string) []ToolCall {
|
|
||||||
var result []ToolCall
|
|
||||||
remaining := text
|
|
||||||
callIdx := 0
|
|
||||||
|
|
||||||
for {
|
|
||||||
_, blockEnd, block, found := findToolCallBlock(remaining)
|
|
||||||
if !found {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
remaining = remaining[blockEnd:]
|
|
||||||
result = append(result, parseInvokeElements(block, &callIdx)...)
|
|
||||||
}
|
|
||||||
|
|
||||||
return result
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseInvokeElements extracts ToolCall entries from <invoke>...</invoke> blocks.
|
|
||||||
func parseInvokeElements(text string, callIdx *int) []ToolCall {
|
|
||||||
var result []ToolCall
|
|
||||||
invokeRemaining := text
|
|
||||||
for {
|
|
||||||
invokeStart := strings.Index(invokeRemaining, "<invoke")
|
|
||||||
if invokeStart == -1 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
invokeEnd := strings.Index(invokeRemaining[invokeStart:], "</invoke>")
|
|
||||||
if invokeEnd == -1 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
invokeBody := invokeRemaining[invokeStart : invokeStart+invokeEnd+len("</invoke>")]
|
|
||||||
invokeRemaining = invokeRemaining[invokeStart+invokeEnd+len("</invoke>"):]
|
|
||||||
|
|
||||||
// Extract tool name from <invoke name="...">
|
|
||||||
nameStart := strings.Index(invokeBody, `name="`)
|
|
||||||
if nameStart == -1 {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
nameStart += len(`name="`)
|
|
||||||
nameEnd := strings.Index(invokeBody[nameStart:], `"`)
|
|
||||||
if nameEnd == -1 {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
toolName := invokeBody[nameStart : nameStart+nameEnd]
|
|
||||||
|
|
||||||
// Extract parameters
|
|
||||||
args := make(map[string]any)
|
|
||||||
paramRemaining := invokeBody
|
|
||||||
for {
|
|
||||||
pStart := strings.Index(paramRemaining, "<parameter")
|
|
||||||
if pStart == -1 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
pNameStart := strings.Index(paramRemaining[pStart:], `name="`)
|
|
||||||
if pNameStart == -1 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
pNameStart += pStart + len(`name="`)
|
|
||||||
pNameEnd := strings.Index(paramRemaining[pNameStart:], `"`)
|
|
||||||
if pNameEnd == -1 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
paramName := paramRemaining[pNameStart : pNameStart+pNameEnd]
|
|
||||||
|
|
||||||
tagClose := strings.Index(paramRemaining[pNameStart:], ">")
|
|
||||||
if tagClose == -1 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
valueStart := pNameStart + tagClose + 1
|
|
||||||
valueEnd := strings.Index(paramRemaining[valueStart:], "</parameter>")
|
|
||||||
if valueEnd == -1 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
paramValue := paramRemaining[valueStart : valueStart+valueEnd]
|
|
||||||
args[paramName] = paramValue
|
|
||||||
paramRemaining = paramRemaining[valueStart+valueEnd+len("</parameter>"):]
|
|
||||||
}
|
|
||||||
|
|
||||||
*callIdx++
|
|
||||||
result = append(result, ToolCall{
|
|
||||||
ID: fmt.Sprintf("xmltc_%d", *callIdx),
|
|
||||||
Type: "function",
|
|
||||||
Name: toolName,
|
|
||||||
Arguments: args,
|
|
||||||
Function: &FunctionCall{
|
|
||||||
Name: toolName,
|
|
||||||
Arguments: cloneToolArgs(args),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}
|
|
||||||
return result
|
|
||||||
}
|
|
||||||
|
|
||||||
// StripXMLToolCalls is the exported version for use by the agent loop.
|
|
||||||
func StripXMLToolCalls(text string) string {
|
|
||||||
return stripXMLToolCalls(text)
|
|
||||||
}
|
|
||||||
|
|
||||||
// stripXMLToolCalls removes XML tool call blocks from response text.
|
|
||||||
// Prevents raw XML tool calls from leaking to users.
|
|
||||||
func stripXMLToolCalls(text string) string {
|
|
||||||
blockStart, blockEnd, _, found := findToolCallBlock(text)
|
|
||||||
if found {
|
|
||||||
cleaned := text[:blockStart] + text[blockEnd:]
|
|
||||||
if _, _, _, more := findToolCallBlock(cleaned); more {
|
|
||||||
cleaned = stripXMLToolCalls(cleaned)
|
|
||||||
}
|
|
||||||
return strings.TrimSpace(cleaned)
|
|
||||||
}
|
|
||||||
|
|
||||||
return strings.TrimSpace(text)
|
|
||||||
}
|
|
||||||
|
|
||||||
// stripToolCallsFromText removes tool call JSON from response text.
|
// stripToolCallsFromText removes tool call JSON from response text.
|
||||||
func stripToolCallsFromText(text string) string {
|
func stripToolCallsFromText(text string) string {
|
||||||
start := strings.Index(text, `{"tool_calls"`)
|
start := strings.Index(text, `{"tool_calls"`)
|
||||||
|
|
|
||||||
262
pkg/providers/tool_call_extract_ext.go
Normal file
262
pkg/providers/tool_call_extract_ext.go
Normal file
|
|
@ -0,0 +1,262 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// --- Shared helpers for XML tool call extraction ---
|
||||||
|
|
||||||
|
// normalizeAlpha keeps only lowercase ASCII letters.
|
||||||
|
// "tool_call" → "toolcall", "Tool-Call" → "toolcall", "ReadFile" → "readfile".
|
||||||
|
func normalizeAlpha(s string) string {
|
||||||
|
var b strings.Builder
|
||||||
|
for _, r := range s {
|
||||||
|
if r >= 'A' && r <= 'Z' {
|
||||||
|
b.WriteRune(r + 32)
|
||||||
|
} else if r >= 'a' && r <= 'z' {
|
||||||
|
b.WriteRune(r)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return b.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
// levenshtein computes the edit distance between two strings.
|
||||||
|
// O(n*m) where n,m are string lengths — negligible for short tag names.
|
||||||
|
func levenshtein(a, b string) int {
|
||||||
|
la, lb := len(a), len(b)
|
||||||
|
if la == 0 {
|
||||||
|
return lb
|
||||||
|
}
|
||||||
|
if lb == 0 {
|
||||||
|
return la
|
||||||
|
}
|
||||||
|
prev := make([]int, lb+1)
|
||||||
|
for j := range prev {
|
||||||
|
prev[j] = j
|
||||||
|
}
|
||||||
|
for i := 1; i <= la; i++ {
|
||||||
|
curr := make([]int, lb+1)
|
||||||
|
curr[0] = i
|
||||||
|
for j := 1; j <= lb; j++ {
|
||||||
|
cost := 1
|
||||||
|
if a[i-1] == b[j-1] {
|
||||||
|
cost = 0
|
||||||
|
}
|
||||||
|
curr[j] = min(curr[j-1]+1, min(prev[j]+1, prev[j-1]+cost))
|
||||||
|
}
|
||||||
|
prev = curr
|
||||||
|
}
|
||||||
|
return prev[lb]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Known tool call tag patterns (already alpha-normalized).
|
||||||
|
// Providers may use different names: tool_call, function_call, tool_use, etc.
|
||||||
|
var toolCallPatterns = []string{"toolcall", "functioncall", "tooluse"}
|
||||||
|
|
||||||
|
// isToolCallTag returns true if the tag name is close to any known tool call
|
||||||
|
// pattern after alpha normalization + edit distance (threshold ≤ 2).
|
||||||
|
func isToolCallTag(name string) bool {
|
||||||
|
const threshold = 2
|
||||||
|
norm := normalizeAlpha(name)
|
||||||
|
for _, pat := range toolCallPatterns {
|
||||||
|
if levenshtein(norm, pat) <= threshold {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// tagSuffix returns the part after the last ':' (namespace separator),
|
||||||
|
// or the whole string if there is no ':'.
|
||||||
|
func tagSuffix(tag string) string {
|
||||||
|
if i := strings.LastIndex(tag, ":"); i >= 0 {
|
||||||
|
return tag[i+1:]
|
||||||
|
}
|
||||||
|
return tag
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- XML block detection via regex ---
|
||||||
|
//
|
||||||
|
// Strategy: find <TAG>…</TAG> pairs using regex, then check if the tag
|
||||||
|
// suffix normalizes to something close to "toolcall" (edit distance ≤ 2).
|
||||||
|
// Uses greedy (longest) match for the closing tag to capture the full block.
|
||||||
|
|
||||||
|
var (
|
||||||
|
reOpenTag = regexp.MustCompile(`<([a-zA-Z][\w:.-]*)>`)
|
||||||
|
reCloseTag = regexp.MustCompile(`</([a-zA-Z][\w:.-]*)>`)
|
||||||
|
reBracketMarker = regexp.MustCompile(`\[TOOLCALL\]`)
|
||||||
|
)
|
||||||
|
|
||||||
|
// findToolCallBlock finds the first XML block whose tag suffix matches
|
||||||
|
// "toolcall" by edit distance. Returns the block boundaries and the inner
|
||||||
|
// content, or found=false.
|
||||||
|
func findToolCallBlock(text string) (blockStart, blockEnd int, content string, found bool) {
|
||||||
|
for _, om := range reOpenTag.FindAllStringSubmatchIndex(text, -1) {
|
||||||
|
tagName := text[om[2]:om[3]]
|
||||||
|
if !isToolCallTag(tagSuffix(tagName)) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Found a toolcall opening tag. Search for the last matching close tag (greedy).
|
||||||
|
afterOpen := text[om[1]:]
|
||||||
|
closes := reCloseTag.FindAllStringSubmatchIndex(afterOpen, -1)
|
||||||
|
for i := len(closes) - 1; i >= 0; i-- {
|
||||||
|
closeTagName := afterOpen[closes[i][2]:closes[i][3]]
|
||||||
|
if isToolCallTag(tagSuffix(closeTagName)) {
|
||||||
|
return om[0], om[1] + closes[i][1], afterOpen[:closes[i][0]], true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback: look for orphaned closing tags (missing opening tag).
|
||||||
|
// Some LLMs emit the closing </ns:tool_call> without a matching opener.
|
||||||
|
// Reconstruct the block start from the first <invoke preceding the closer.
|
||||||
|
for _, cm := range reCloseTag.FindAllStringSubmatchIndex(text, -1) {
|
||||||
|
closeTagName := text[cm[2]:cm[3]]
|
||||||
|
if !isToolCallTag(tagSuffix(closeTagName)) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Found an orphaned toolcall closing tag. Scan backwards for <invoke.
|
||||||
|
before := text[:cm[0]]
|
||||||
|
invokePos := strings.LastIndex(before, "<invoke")
|
||||||
|
if invokePos == -1 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Also consume a preceding [TOOLCALL] marker if present.
|
||||||
|
start := invokePos
|
||||||
|
if loc := reBracketMarker.FindStringIndex(before[:start]); loc != nil &&
|
||||||
|
strings.TrimSpace(before[loc[1]:start]) == "" {
|
||||||
|
start = loc[0]
|
||||||
|
}
|
||||||
|
return start, cm[1], text[invokePos:cm[0]], true
|
||||||
|
}
|
||||||
|
|
||||||
|
return 0, 0, "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExtractXMLToolCalls extracts tool calls from XML-formatted text.
|
||||||
|
//
|
||||||
|
// Expected format:
|
||||||
|
//
|
||||||
|
// <ns:toolcall>
|
||||||
|
// <invoke name="tool_name">
|
||||||
|
// <parameter name="param">value</parameter>
|
||||||
|
// </invoke>
|
||||||
|
// </ns:toolcall>
|
||||||
|
func ExtractXMLToolCalls(text string) []ToolCall {
|
||||||
|
return extractXMLToolCalls(text)
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractXMLToolCalls(text string) []ToolCall {
|
||||||
|
var result []ToolCall
|
||||||
|
remaining := text
|
||||||
|
callIdx := 0
|
||||||
|
|
||||||
|
for {
|
||||||
|
_, blockEnd, block, found := findToolCallBlock(remaining)
|
||||||
|
if !found {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
remaining = remaining[blockEnd:]
|
||||||
|
result = append(result, parseInvokeElements(block, &callIdx)...)
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseInvokeElements extracts ToolCall entries from <invoke>...</invoke> blocks.
|
||||||
|
func parseInvokeElements(text string, callIdx *int) []ToolCall {
|
||||||
|
var result []ToolCall
|
||||||
|
invokeRemaining := text
|
||||||
|
for {
|
||||||
|
invokeStart := strings.Index(invokeRemaining, "<invoke")
|
||||||
|
if invokeStart == -1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
invokeEnd := strings.Index(invokeRemaining[invokeStart:], "</invoke>")
|
||||||
|
if invokeEnd == -1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
invokeBody := invokeRemaining[invokeStart : invokeStart+invokeEnd+len("</invoke>")]
|
||||||
|
invokeRemaining = invokeRemaining[invokeStart+invokeEnd+len("</invoke>"):]
|
||||||
|
|
||||||
|
// Extract tool name from <invoke name="...">
|
||||||
|
nameStart := strings.Index(invokeBody, `name="`)
|
||||||
|
if nameStart == -1 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
nameStart += len(`name="`)
|
||||||
|
nameEnd := strings.Index(invokeBody[nameStart:], `"`)
|
||||||
|
if nameEnd == -1 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
toolName := invokeBody[nameStart : nameStart+nameEnd]
|
||||||
|
|
||||||
|
// Extract parameters
|
||||||
|
args := make(map[string]any)
|
||||||
|
paramRemaining := invokeBody
|
||||||
|
for {
|
||||||
|
pStart := strings.Index(paramRemaining, "<parameter")
|
||||||
|
if pStart == -1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
pNameStart := strings.Index(paramRemaining[pStart:], `name="`)
|
||||||
|
if pNameStart == -1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
pNameStart += pStart + len(`name="`)
|
||||||
|
pNameEnd := strings.Index(paramRemaining[pNameStart:], `"`)
|
||||||
|
if pNameEnd == -1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
paramName := paramRemaining[pNameStart : pNameStart+pNameEnd]
|
||||||
|
|
||||||
|
tagClose := strings.Index(paramRemaining[pNameStart:], ">")
|
||||||
|
if tagClose == -1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
valueStart := pNameStart + tagClose + 1
|
||||||
|
valueEnd := strings.Index(paramRemaining[valueStart:], "</parameter>")
|
||||||
|
if valueEnd == -1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
paramValue := paramRemaining[valueStart : valueStart+valueEnd]
|
||||||
|
args[paramName] = paramValue
|
||||||
|
paramRemaining = paramRemaining[valueStart+valueEnd+len("</parameter>"):]
|
||||||
|
}
|
||||||
|
|
||||||
|
*callIdx++
|
||||||
|
result = append(result, ToolCall{
|
||||||
|
ID: fmt.Sprintf("xmltc_%d", *callIdx),
|
||||||
|
Type: "function",
|
||||||
|
Name: toolName,
|
||||||
|
Arguments: args,
|
||||||
|
Function: &FunctionCall{
|
||||||
|
Name: toolName,
|
||||||
|
Arguments: cloneToolArgs(args),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// StripXMLToolCalls is the exported version for use by the agent loop.
|
||||||
|
func StripXMLToolCalls(text string) string {
|
||||||
|
return stripXMLToolCalls(text)
|
||||||
|
}
|
||||||
|
|
||||||
|
// stripXMLToolCalls removes XML tool call blocks from response text.
|
||||||
|
// Prevents raw XML tool calls from leaking to users.
|
||||||
|
func stripXMLToolCalls(text string) string {
|
||||||
|
blockStart, blockEnd, _, found := findToolCallBlock(text)
|
||||||
|
if found {
|
||||||
|
cleaned := text[:blockStart] + text[blockEnd:]
|
||||||
|
if _, _, _, more := findToolCallBlock(cleaned); more {
|
||||||
|
cleaned = stripXMLToolCalls(cleaned)
|
||||||
|
}
|
||||||
|
return strings.TrimSpace(cleaned)
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.TrimSpace(text)
|
||||||
|
}
|
||||||
|
|
@ -190,8 +190,3 @@ func resolveLinkedPeerID(identityLinks map[string][]string, channel, peerID stri
|
||||||
}
|
}
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
// BuildSubagentSessionKey returns "subagent:<taskID>" for subagent sessions.
|
|
||||||
func BuildSubagentSessionKey(taskID string) string {
|
|
||||||
return fmt.Sprintf("subagent:%s", taskID)
|
|
||||||
}
|
|
||||||
|
|
|
||||||
8
pkg/routing/session_key_ext.go
Normal file
8
pkg/routing/session_key_ext.go
Normal file
|
|
@ -0,0 +1,8 @@
|
||||||
|
package routing
|
||||||
|
|
||||||
|
import "fmt"
|
||||||
|
|
||||||
|
// BuildSubagentSessionKey returns "subagent:<taskID>" for subagent sessions.
|
||||||
|
func BuildSubagentSessionKey(taskID string) string {
|
||||||
|
return fmt.Sprintf("subagent:%s", taskID)
|
||||||
|
}
|
||||||
|
|
@ -4,11 +4,11 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/constants"
|
||||||
"github.com/sipeed/picoclaw/pkg/cron"
|
"github.com/sipeed/picoclaw/pkg/cron"
|
||||||
"github.com/sipeed/picoclaw/pkg/utils"
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
)
|
)
|
||||||
|
|
@ -21,18 +21,9 @@ type JobExecutor interface {
|
||||||
// CronTool provides scheduling capabilities for the agent
|
// CronTool provides scheduling capabilities for the agent
|
||||||
type CronTool struct {
|
type CronTool struct {
|
||||||
cronService *cron.CronService
|
cronService *cron.CronService
|
||||||
|
|
||||||
executor JobExecutor
|
executor JobExecutor
|
||||||
|
|
||||||
msgBus *bus.MessageBus
|
msgBus *bus.MessageBus
|
||||||
|
|
||||||
execTool *ExecTool
|
execTool *ExecTool
|
||||||
|
|
||||||
channel string
|
|
||||||
|
|
||||||
chatID string
|
|
||||||
|
|
||||||
mu sync.RWMutex
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewCronTool creates a new CronTool
|
// NewCronTool creates a new CronTool
|
||||||
|
|
@ -49,11 +40,8 @@ func NewCronTool(
|
||||||
execTool.SetTimeout(execTimeout)
|
execTool.SetTimeout(execTimeout)
|
||||||
return &CronTool{
|
return &CronTool{
|
||||||
cronService: cronService,
|
cronService: cronService,
|
||||||
|
|
||||||
executor: executor,
|
executor: executor,
|
||||||
|
|
||||||
msgBus: msgBus,
|
msgBus: msgBus,
|
||||||
|
|
||||||
execTool: execTool,
|
execTool: execTool,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
@ -75,45 +63,39 @@ func (t *CronTool) Parameters() map[string]any {
|
||||||
"properties": map[string]any{
|
"properties": map[string]any{
|
||||||
"action": map[string]any{
|
"action": map[string]any{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
|
|
||||||
"enum": []string{"add", "list", "remove", "enable", "disable"},
|
"enum": []string{"add", "list", "remove", "enable", "disable"},
|
||||||
|
|
||||||
"description": "Action to perform. Use 'add' when user wants to schedule a reminder or task.",
|
"description": "Action to perform. Use 'add' when user wants to schedule a reminder or task.",
|
||||||
},
|
},
|
||||||
"message": map[string]any{
|
"message": map[string]any{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
|
|
||||||
"description": "The reminder/task message to display when triggered. If 'command' is used, this describes what the command does.",
|
"description": "The reminder/task message to display when triggered. If 'command' is used, this describes what the command does.",
|
||||||
},
|
},
|
||||||
"command": map[string]any{
|
"command": map[string]any{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
|
|
||||||
"description": "Optional: Shell command to execute directly (e.g., 'df -h'). If set, the agent will run this command and report output instead of just showing the message. 'deliver' will be forced to false for commands.",
|
"description": "Optional: Shell command to execute directly (e.g., 'df -h'). If set, the agent will run this command and report output instead of just showing the message. 'deliver' will be forced to false for commands.",
|
||||||
},
|
},
|
||||||
|
"command_confirm": map[string]any{
|
||||||
|
"type": "boolean",
|
||||||
|
"description": "Required when using command=true. Must be true to explicitly confirm scheduling a shell command.",
|
||||||
|
},
|
||||||
"at_seconds": map[string]any{
|
"at_seconds": map[string]any{
|
||||||
"type": "integer",
|
"type": "integer",
|
||||||
|
|
||||||
"description": "One-time reminder: seconds from now when to trigger (e.g., 600 for 10 minutes later). Use this for one-time reminders like 'remind me in 10 minutes'.",
|
"description": "One-time reminder: seconds from now when to trigger (e.g., 600 for 10 minutes later). Use this for one-time reminders like 'remind me in 10 minutes'.",
|
||||||
},
|
},
|
||||||
"every_seconds": map[string]any{
|
"every_seconds": map[string]any{
|
||||||
"type": "integer",
|
"type": "integer",
|
||||||
|
|
||||||
"description": "Recurring interval in seconds (e.g., 3600 for every hour). Use this ONLY for recurring tasks like 'every 2 hours' or 'daily reminder'.",
|
"description": "Recurring interval in seconds (e.g., 3600 for every hour). Use this ONLY for recurring tasks like 'every 2 hours' or 'daily reminder'.",
|
||||||
},
|
},
|
||||||
"cron_expr": map[string]any{
|
"cron_expr": map[string]any{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
|
|
||||||
"description": "Cron expression for complex recurring schedules (e.g., '0 9 * * *' for daily at 9am). Use this for complex recurring schedules.",
|
"description": "Cron expression for complex recurring schedules (e.g., '0 9 * * *' for daily at 9am). Use this for complex recurring schedules.",
|
||||||
},
|
},
|
||||||
"job_id": map[string]any{
|
"job_id": map[string]any{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
|
|
||||||
"description": "Job ID (for remove/enable/disable)",
|
"description": "Job ID (for remove/enable/disable)",
|
||||||
},
|
},
|
||||||
"deliver": map[string]any{
|
"deliver": map[string]any{
|
||||||
"type": "boolean",
|
"type": "boolean",
|
||||||
|
|
||||||
"description": "If true, send message directly to channel. If false, let agent process message (for complex tasks). Default: true",
|
"description": "If true, send message directly to channel. If false, let agent process message (for complex tasks). Default: true",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
|
|
@ -121,18 +103,6 @@ func (t *CronTool) Parameters() map[string]any {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetContext sets the current session context for job creation
|
|
||||||
|
|
||||||
func (t *CronTool) SetContext(channel, chatID string) {
|
|
||||||
t.mu.Lock()
|
|
||||||
|
|
||||||
defer t.mu.Unlock()
|
|
||||||
|
|
||||||
t.channel = channel
|
|
||||||
|
|
||||||
t.chatID = chatID
|
|
||||||
}
|
|
||||||
|
|
||||||
// Execute runs the tool with the given arguments
|
// Execute runs the tool with the given arguments
|
||||||
func (t *CronTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
func (t *CronTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
action, ok := args["action"].(string)
|
action, ok := args["action"].(string)
|
||||||
|
|
@ -142,9 +112,7 @@ func (t *CronTool) Execute(ctx context.Context, args map[string]any) *ToolResult
|
||||||
|
|
||||||
switch action {
|
switch action {
|
||||||
case "add":
|
case "add":
|
||||||
|
return t.addJob(ctx, args)
|
||||||
return t.addJob(args)
|
|
||||||
|
|
||||||
case "list":
|
case "list":
|
||||||
return t.listJobs()
|
return t.listJobs()
|
||||||
case "remove":
|
case "remove":
|
||||||
|
|
@ -158,14 +126,9 @@ func (t *CronTool) Execute(ctx context.Context, args map[string]any) *ToolResult
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *CronTool) addJob(args map[string]any) *ToolResult {
|
func (t *CronTool) addJob(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
t.mu.RLock()
|
channel := ToolChannel(ctx)
|
||||||
|
chatID := ToolChatID(ctx)
|
||||||
channel := t.channel
|
|
||||||
|
|
||||||
chatID := t.chatID
|
|
||||||
|
|
||||||
t.mu.RUnlock()
|
|
||||||
|
|
||||||
if channel == "" || chatID == "" {
|
if channel == "" || chatID == "" {
|
||||||
return ErrorResult("no session context (channel/chat_id not set). Use this tool in an active conversation.")
|
return ErrorResult("no session context (channel/chat_id not set). Use this tool in an active conversation.")
|
||||||
|
|
@ -183,6 +146,12 @@ func (t *CronTool) addJob(args map[string]any) *ToolResult {
|
||||||
everySeconds, hasEvery := args["every_seconds"].(float64)
|
everySeconds, hasEvery := args["every_seconds"].(float64)
|
||||||
cronExpr, hasCron := args["cron_expr"].(string)
|
cronExpr, hasCron := args["cron_expr"].(string)
|
||||||
|
|
||||||
|
// Fix: type assertions return true for zero values, need additional validity checks
|
||||||
|
// This prevents LLMs that fill unused optional parameters with defaults (0) from triggering wrong type
|
||||||
|
hasAt = hasAt && atSeconds > 0
|
||||||
|
hasEvery = hasEvery && everySeconds > 0
|
||||||
|
hasCron = hasCron && cronExpr != ""
|
||||||
|
|
||||||
// Priority: at_seconds > every_seconds > cron_expr
|
// Priority: at_seconds > every_seconds > cron_expr
|
||||||
if hasAt {
|
if hasAt {
|
||||||
atMS := time.Now().UnixMilli() + int64(atSeconds)*1000
|
atMS := time.Now().UnixMilli() + int64(atSeconds)*1000
|
||||||
|
|
@ -194,7 +163,6 @@ func (t *CronTool) addJob(args map[string]any) *ToolResult {
|
||||||
everyMS := int64(everySeconds) * 1000
|
everyMS := int64(everySeconds) * 1000
|
||||||
schedule = cron.CronSchedule{
|
schedule = cron.CronSchedule{
|
||||||
Kind: "every",
|
Kind: "every",
|
||||||
|
|
||||||
EveryMS: &everyMS,
|
EveryMS: &everyMS,
|
||||||
}
|
}
|
||||||
} else if hasCron {
|
} else if hasCron {
|
||||||
|
|
@ -212,17 +180,17 @@ func (t *CronTool) addJob(args map[string]any) *ToolResult {
|
||||||
deliver = d
|
deliver = d
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GHSA-pv8c-p6jf-3fpp: command scheduling requires internal channel + explicit confirm.
|
||||||
|
// Non-command reminders (plain messages) remain open to all channels.
|
||||||
command, _ := args["command"].(string)
|
command, _ := args["command"].(string)
|
||||||
|
commandConfirm, _ := args["command_confirm"].(bool)
|
||||||
if command != "" {
|
if command != "" {
|
||||||
// Commands must be processed by agent/exec tool, so deliver must be false (or handled specifically)
|
if !constants.IsInternalChannel(channel) {
|
||||||
|
return ErrorResult("scheduling command execution is restricted to internal channels")
|
||||||
// Actually, let's keep deliver=false to let the system know it's not a simple chat message
|
}
|
||||||
|
if !commandConfirm {
|
||||||
// But for our new logic in ExecuteJob, we can handle it regardless of deliver flag if Payload.Command is set.
|
return ErrorResult("command_confirm=true is required to schedule command execution")
|
||||||
|
}
|
||||||
// However, logically, it's not "delivered" to chat directly as is.
|
|
||||||
|
|
||||||
deliver = false
|
deliver = false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -257,10 +225,8 @@ func (t *CronTool) listJobs() *ToolResult {
|
||||||
return SilentResult("No scheduled jobs")
|
return SilentResult("No scheduled jobs")
|
||||||
}
|
}
|
||||||
|
|
||||||
var sb strings.Builder
|
var result strings.Builder
|
||||||
|
result.WriteString("Scheduled jobs:\n")
|
||||||
sb.WriteString("Scheduled jobs:\n")
|
|
||||||
|
|
||||||
for _, j := range jobs {
|
for _, j := range jobs {
|
||||||
var scheduleInfo string
|
var scheduleInfo string
|
||||||
if j.Schedule.Kind == "every" && j.Schedule.EveryMS != nil {
|
if j.Schedule.Kind == "every" && j.Schedule.EveryMS != nil {
|
||||||
|
|
@ -272,11 +238,10 @@ func (t *CronTool) listJobs() *ToolResult {
|
||||||
} else {
|
} else {
|
||||||
scheduleInfo = "unknown"
|
scheduleInfo = "unknown"
|
||||||
}
|
}
|
||||||
|
result.WriteString(fmt.Sprintf("- %s (id: %s, %s)\n", j.Name, j.ID, scheduleInfo))
|
||||||
fmt.Fprintf(&sb, "- %s (id: %s, %s)\n", j.Name, j.ID, scheduleInfo)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return SilentResult(sb.String())
|
return SilentResult(result.String())
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *CronTool) removeJob(args map[string]any) *ToolResult {
|
func (t *CronTool) removeJob(args map[string]any) *ToolResult {
|
||||||
|
|
@ -327,6 +292,8 @@ func (t *CronTool) ExecuteJob(ctx context.Context, job *cron.CronJob) string {
|
||||||
if job.Payload.Command != "" {
|
if job.Payload.Command != "" {
|
||||||
args := map[string]any{
|
args := map[string]any{
|
||||||
"command": job.Payload.Command,
|
"command": job.Payload.Command,
|
||||||
|
"__channel": channel,
|
||||||
|
"__chat_id": chatID,
|
||||||
}
|
}
|
||||||
|
|
||||||
result := t.execTool.Execute(ctx, args)
|
result := t.execTool.Execute(ctx, args)
|
||||||
|
|
@ -341,9 +308,7 @@ func (t *CronTool) ExecuteJob(ctx context.Context, job *cron.CronJob) string {
|
||||||
defer pubCancel()
|
defer pubCancel()
|
||||||
t.msgBus.PublishOutbound(pubCtx, bus.OutboundMessage{
|
t.msgBus.PublishOutbound(pubCtx, bus.OutboundMessage{
|
||||||
Channel: channel,
|
Channel: channel,
|
||||||
|
|
||||||
ChatID: chatID,
|
ChatID: chatID,
|
||||||
|
|
||||||
Content: output,
|
Content: output,
|
||||||
})
|
})
|
||||||
return "ok"
|
return "ok"
|
||||||
|
|
@ -355,9 +320,7 @@ func (t *CronTool) ExecuteJob(ctx context.Context, job *cron.CronJob) string {
|
||||||
defer pubCancel()
|
defer pubCancel()
|
||||||
t.msgBus.PublishOutbound(pubCtx, bus.OutboundMessage{
|
t.msgBus.PublishOutbound(pubCtx, bus.OutboundMessage{
|
||||||
Channel: channel,
|
Channel: channel,
|
||||||
|
|
||||||
ChatID: chatID,
|
ChatID: chatID,
|
||||||
|
|
||||||
Content: job.Payload.Message,
|
Content: job.Payload.Message,
|
||||||
})
|
})
|
||||||
return "ok"
|
return "ok"
|
||||||
|
|
|
||||||
|
|
@ -1,98 +1,17 @@
|
||||||
package utils
|
package utils
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"regexp"
|
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
"unicode"
|
"unicode"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Repetition detection constants.
|
// Global variable to disable truncation
|
||||||
const (
|
var disableTruncation atomic.Bool
|
||||||
repetitionSampleSize = 2000 // runes to sample from the tail
|
|
||||||
repetitionNgramSize = 10 // sliding window length
|
|
||||||
repetitionUniqueThreshold = 0.1 // unique ratio below this → repetition
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
// SetDisableTruncation globally enables or disables string truncation
|
||||||
thinkBlockClosedRe = regexp.MustCompile(`(?is)<think>.*?</think>`)
|
func SetDisableTruncation(enabled bool) {
|
||||||
thinkBlockOpenRe = regexp.MustCompile(`(?is)<think>.*$`)
|
disableTruncation.Store(enabled)
|
||||||
)
|
|
||||||
|
|
||||||
// StripThinkBlocks removes <think>…</think> blocks (including unclosed ones)
|
|
||||||
// from s and returns the trimmed result.
|
|
||||||
func StripThinkBlocks(s string) string {
|
|
||||||
s = thinkBlockClosedRe.ReplaceAllString(s, "")
|
|
||||||
s = thinkBlockOpenRe.ReplaceAllString(s, "")
|
|
||||||
return strings.TrimSpace(s)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TailPad returns a fixed-height block of n visual lines built from the
|
|
||||||
// tail of s. Long lines are wrapped at wrapWidth runes so the result
|
|
||||||
// never exceeds the chat bubble width. If fewer than n visual lines
|
|
||||||
// exist, Braille-blank lines (\u2800) are prepended as padding.
|
|
||||||
func TailPad(s string, n, wrapWidth int) string {
|
|
||||||
// Wrap each raw line into visual lines respecting wrapWidth.
|
|
||||||
var visual []string
|
|
||||||
for _, raw := range strings.Split(s, "\n") {
|
|
||||||
visual = append(visual, wrapLine(raw, wrapWidth)...)
|
|
||||||
}
|
|
||||||
if len(visual) > n {
|
|
||||||
visual = visual[len(visual)-n:]
|
|
||||||
}
|
|
||||||
for len(visual) < n {
|
|
||||||
visual = append([]string{"\u2800"}, visual...)
|
|
||||||
}
|
|
||||||
return strings.Join(visual, "\n")
|
|
||||||
}
|
|
||||||
|
|
||||||
// wrapLine splits a single line into segments of at most width runes.
|
|
||||||
// An empty line produces one empty string (preserving blank lines).
|
|
||||||
func wrapLine(line string, width int) []string {
|
|
||||||
// ASCII fast path: byte length == rune length for pure ASCII
|
|
||||||
if len(line) <= width {
|
|
||||||
return []string{line}
|
|
||||||
}
|
|
||||||
runes := []rune(line)
|
|
||||||
if len(runes) <= width {
|
|
||||||
return []string{line}
|
|
||||||
}
|
|
||||||
var segs []string
|
|
||||||
for len(runes) > 0 {
|
|
||||||
end := width
|
|
||||||
if end > len(runes) {
|
|
||||||
end = len(runes)
|
|
||||||
}
|
|
||||||
segs = append(segs, string(runes[:end]))
|
|
||||||
runes = runes[end:]
|
|
||||||
}
|
|
||||||
return segs
|
|
||||||
}
|
|
||||||
|
|
||||||
// DetectRepetitionLoop checks if text contains degenerate repetition
|
|
||||||
// by computing the unique N-gram ratio on the last repetitionSampleSize runes.
|
|
||||||
// Returns true if the ratio of unique N-grams to total N-grams
|
|
||||||
// falls below repetitionUniqueThreshold (i.e., 90%+ are duplicates).
|
|
||||||
func DetectRepetitionLoop(text string) bool {
|
|
||||||
runes := []rune(text)
|
|
||||||
|
|
||||||
// Sample the tail
|
|
||||||
if len(runes) > repetitionSampleSize {
|
|
||||||
runes = runes[len(runes)-repetitionSampleSize:]
|
|
||||||
}
|
|
||||||
|
|
||||||
total := len(runes) - repetitionNgramSize + 1
|
|
||||||
if total <= 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
unique := make(map[string]struct{}, total/repetitionNgramSize)
|
|
||||||
for i := 0; i < total; i++ {
|
|
||||||
ng := string(runes[i : i+repetitionNgramSize])
|
|
||||||
unique[ng] = struct{}{}
|
|
||||||
}
|
|
||||||
|
|
||||||
ratio := float64(len(unique)) / float64(total)
|
|
||||||
return ratio < repetitionUniqueThreshold
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// SanitizeMessageContent removes Unicode control characters, format characters (RTL overrides,
|
// SanitizeMessageContent removes Unicode control characters, format characters (RTL overrides,
|
||||||
|
|
@ -100,9 +19,14 @@ func DetectRepetitionLoop(text string) bool {
|
||||||
// or cause display issues in the agent UI.
|
// or cause display issues in the agent UI.
|
||||||
func SanitizeMessageContent(input string) string {
|
func SanitizeMessageContent(input string) string {
|
||||||
var sb strings.Builder
|
var sb strings.Builder
|
||||||
|
// Pre-allocate memory to avoid multiple allocations
|
||||||
sb.Grow(len(input))
|
sb.Grow(len(input))
|
||||||
|
|
||||||
for _, r := range input {
|
for _, r := range input {
|
||||||
|
// unicode.IsGraphic returns true if the rune is a Unicode graphic character.
|
||||||
|
// This includes letters, marks, numbers, punctuation, and symbols.
|
||||||
|
// It excludes control characters (Cc), format characters (Cf),
|
||||||
|
// surrogates (Cs), and private use (Co).
|
||||||
if unicode.IsGraphic(r) || r == '\n' || r == '\r' || r == '\t' {
|
if unicode.IsGraphic(r) || r == '\n' || r == '\r' || r == '\t' {
|
||||||
sb.WriteRune(r)
|
sb.WriteRune(r)
|
||||||
}
|
}
|
||||||
|
|
@ -115,13 +39,13 @@ func SanitizeMessageContent(input string) string {
|
||||||
// Handles multi-byte Unicode characters properly.
|
// Handles multi-byte Unicode characters properly.
|
||||||
// If the string is truncated, "..." is appended to indicate truncation.
|
// If the string is truncated, "..." is appended to indicate truncation.
|
||||||
func Truncate(s string, maxLen int) string {
|
func Truncate(s string, maxLen int) string {
|
||||||
|
// If the no-truncate flag is active, it returns the full string
|
||||||
|
if disableTruncation.Load() {
|
||||||
|
return s
|
||||||
|
}
|
||||||
if maxLen <= 0 {
|
if maxLen <= 0 {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
// ASCII fast path: byte length == rune length for pure ASCII
|
|
||||||
if len(s) <= maxLen {
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
runes := []rune(s)
|
runes := []rune(s)
|
||||||
if len(runes) <= maxLen {
|
if len(runes) <= maxLen {
|
||||||
return s
|
return s
|
||||||
|
|
|
||||||
95
pkg/utils/string_ext.go
Normal file
95
pkg/utils/string_ext.go
Normal file
|
|
@ -0,0 +1,95 @@
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Repetition detection constants.
|
||||||
|
const (
|
||||||
|
repetitionSampleSize = 2000 // runes to sample from the tail
|
||||||
|
repetitionNgramSize = 10 // sliding window length
|
||||||
|
repetitionUniqueThreshold = 0.1 // unique ratio below this → repetition
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
thinkBlockClosedRe = regexp.MustCompile(`(?is)<think>.*?</think>`)
|
||||||
|
thinkBlockOpenRe = regexp.MustCompile(`(?is)<think>.*$`)
|
||||||
|
)
|
||||||
|
|
||||||
|
// StripThinkBlocks removes <think>…</think> blocks (including unclosed ones)
|
||||||
|
// from s and returns the trimmed result.
|
||||||
|
func StripThinkBlocks(s string) string {
|
||||||
|
s = thinkBlockClosedRe.ReplaceAllString(s, "")
|
||||||
|
s = thinkBlockOpenRe.ReplaceAllString(s, "")
|
||||||
|
return strings.TrimSpace(s)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TailPad returns a fixed-height block of n visual lines built from the
|
||||||
|
// tail of s. Long lines are wrapped at wrapWidth runes so the result
|
||||||
|
// never exceeds the chat bubble width. If fewer than n visual lines
|
||||||
|
// exist, Braille-blank lines (\u2800) are prepended as padding.
|
||||||
|
func TailPad(s string, n, wrapWidth int) string {
|
||||||
|
// Wrap each raw line into visual lines respecting wrapWidth.
|
||||||
|
var visual []string
|
||||||
|
for _, raw := range strings.Split(s, "\n") {
|
||||||
|
visual = append(visual, wrapLine(raw, wrapWidth)...)
|
||||||
|
}
|
||||||
|
if len(visual) > n {
|
||||||
|
visual = visual[len(visual)-n:]
|
||||||
|
}
|
||||||
|
for len(visual) < n {
|
||||||
|
visual = append([]string{"\u2800"}, visual...)
|
||||||
|
}
|
||||||
|
return strings.Join(visual, "\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
// wrapLine splits a single line into segments of at most width runes.
|
||||||
|
// An empty line produces one empty string (preserving blank lines).
|
||||||
|
func wrapLine(line string, width int) []string {
|
||||||
|
// ASCII fast path: byte length == rune length for pure ASCII
|
||||||
|
if len(line) <= width {
|
||||||
|
return []string{line}
|
||||||
|
}
|
||||||
|
runes := []rune(line)
|
||||||
|
if len(runes) <= width {
|
||||||
|
return []string{line}
|
||||||
|
}
|
||||||
|
var segs []string
|
||||||
|
for len(runes) > 0 {
|
||||||
|
end := width
|
||||||
|
if end > len(runes) {
|
||||||
|
end = len(runes)
|
||||||
|
}
|
||||||
|
segs = append(segs, string(runes[:end]))
|
||||||
|
runes = runes[end:]
|
||||||
|
}
|
||||||
|
return segs
|
||||||
|
}
|
||||||
|
|
||||||
|
// DetectRepetitionLoop checks if text contains degenerate repetition
|
||||||
|
// by computing the unique N-gram ratio on the last repetitionSampleSize runes.
|
||||||
|
// Returns true if the ratio of unique N-grams to total N-grams
|
||||||
|
// falls below repetitionUniqueThreshold (i.e., 90%+ are duplicates).
|
||||||
|
func DetectRepetitionLoop(text string) bool {
|
||||||
|
runes := []rune(text)
|
||||||
|
|
||||||
|
// Sample the tail
|
||||||
|
if len(runes) > repetitionSampleSize {
|
||||||
|
runes = runes[len(runes)-repetitionSampleSize:]
|
||||||
|
}
|
||||||
|
|
||||||
|
total := len(runes) - repetitionNgramSize + 1
|
||||||
|
if total <= 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
unique := make(map[string]struct{}, total/repetitionNgramSize)
|
||||||
|
for i := 0; i < total; i++ {
|
||||||
|
ng := string(runes[i : i+repetitionNgramSize])
|
||||||
|
unique[ng] = struct{}{}
|
||||||
|
}
|
||||||
|
|
||||||
|
ratio := float64(len(unique)) / float64(total)
|
||||||
|
return ratio < repetitionUniqueThreshold
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue