Merge branch 'huaaudio-feat/audio-call'
This commit is contained in:
commit
bf82c14842
16 changed files with 1139 additions and 184 deletions
|
|
@ -61,6 +61,9 @@ linters:
|
|||
- usestdlibvars
|
||||
- usetesting
|
||||
settings:
|
||||
gomoddirectives:
|
||||
replace-allow-list:
|
||||
- github.com/bwmarrin/discordgo
|
||||
errcheck:
|
||||
check-type-assertions: true
|
||||
check-blank: true
|
||||
|
|
|
|||
6
go.mod
6
go.mod
|
|
@ -22,6 +22,8 @@ require (
|
|||
github.com/mymmrac/telego v1.7.0
|
||||
github.com/open-dingtalk/dingtalk-stream-sdk-go v0.9.1
|
||||
github.com/openai/openai-go/v3 v3.22.0
|
||||
github.com/pion/rtp v1.8.7
|
||||
github.com/pion/webrtc/v3 v3.3.6
|
||||
github.com/rivo/tview v0.42.0
|
||||
github.com/rs/zerolog v1.34.0
|
||||
github.com/slack-go/slack v0.17.3
|
||||
|
|
@ -41,6 +43,7 @@ require (
|
|||
require (
|
||||
filippo.io/edwards25519 v1.2.0 // indirect
|
||||
github.com/beeper/argo-go v1.1.2 // indirect
|
||||
github.com/cloudflare/circl v1.6.3 // indirect
|
||||
github.com/coder/websocket v1.8.14 // indirect
|
||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||
|
|
@ -53,6 +56,7 @@ require (
|
|||
github.com/mattn/go-isatty v0.0.20 // indirect
|
||||
github.com/ncruces/go-strftime v1.0.0 // indirect
|
||||
github.com/petermattis/goid v0.0.0-20260226131333-17d1149c6ac6 // indirect
|
||||
github.com/pion/randutil v0.1.0 // indirect
|
||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||
github.com/rivo/uniseg v0.4.7 // indirect
|
||||
|
|
@ -98,3 +102,5 @@ require (
|
|||
golang.org/x/sync v0.20.0 // indirect
|
||||
golang.org/x/sys v0.42.0 // indirect
|
||||
)
|
||||
|
||||
replace github.com/bwmarrin/discordgo => github.com/yeongaori/discordgo-fork v0.0.0-20260319072544-e8e546f5d532
|
||||
|
|
|
|||
13
go.sum
13
go.sum
|
|
@ -19,8 +19,6 @@ github.com/anthropics/anthropic-sdk-go v1.26.0 h1:oUTzFaUpAevfuELAP1sjL6CQJ9HHAf
|
|||
github.com/anthropics/anthropic-sdk-go v1.26.0/go.mod h1:qUKmaW+uuPB64iy1l+4kOSvaLqPXnHTTBKH6RVZ7q5Q=
|
||||
github.com/beeper/argo-go v1.1.2 h1:UQI2G8F+NLfGTOmTUI0254pGKx/HUU/etbUGTJv91Fs=
|
||||
github.com/beeper/argo-go v1.1.2/go.mod h1:M+LJAnyowKVQ6Rdj6XYGEn+qcVFkb3R/MUpqkGR0hM4=
|
||||
github.com/bwmarrin/discordgo v0.29.0 h1:FmWeXFaKUwrcL3Cx65c20bTRW+vOb6k8AnaP+EgjDno=
|
||||
github.com/bwmarrin/discordgo v0.29.0/go.mod h1:NJZpH+1AfhIcyQsPeuBKsUtYrRnjkyu0kIVMCHkZtRY=
|
||||
github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M=
|
||||
github.com/bytedance/gopkg v0.1.3/go.mod h1:576VvJ+eJgyCzdjS+c4+77QF3p7ubbtiKARP3TxducM=
|
||||
github.com/bytedance/sonic v1.15.0 h1:/PXeWFaR5ElNcVE84U0dOHjiMHQOwNIx3K4ymzh/uSE=
|
||||
|
|
@ -31,6 +29,8 @@ github.com/caarlos0/env/v11 v11.4.0 h1:Kcb6t5kIIr4XkoQC9AF2j+8E1Jsrl3Wz/hhm1LtoG
|
|||
github.com/caarlos0/env/v11 v11.4.0/go.mod h1:qupehSf/Y0TUTsxKywqRt/vJjN5nz6vauiYEUUr8P4U=
|
||||
github.com/cespare/xxhash/v2 v2.1.2/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||
github.com/cespare/xxhash/v2 v2.2.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||
github.com/cloudflare/circl v1.6.3 h1:9GPOhQGF9MCYUeXyMYlqTR6a5gTrgR/fBLXvUgtVcg8=
|
||||
github.com/cloudflare/circl v1.6.3/go.mod h1:2eXP6Qfat4O/Yhh8BznvKnJ+uzEoTQ6jVKJRn81BiS4=
|
||||
github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M=
|
||||
github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU=
|
||||
github.com/coder/websocket v1.8.14 h1:9L0p0iKiNOibykf283eHkKUHHrpG7f65OE3BhhO7v9g=
|
||||
|
|
@ -162,6 +162,12 @@ github.com/openai/openai-go/v3 v3.22.0 h1:6MEoNoV8sbjOVmXdvhmuX3BjVbVdcExbVyGixi
|
|||
github.com/openai/openai-go/v3 v3.22.0/go.mod h1:cdufnVK14cWcT9qA1rRtrXx4FTRsgbDPW7Ia7SS5cZo=
|
||||
github.com/petermattis/goid v0.0.0-20260226131333-17d1149c6ac6 h1:rh2lKw/P/EqHa724vYH2+VVQ1YnW4u6EOXl0PMAovZE=
|
||||
github.com/petermattis/goid v0.0.0-20260226131333-17d1149c6ac6/go.mod h1:pxMtw7cyUw6B2bRH0ZBANSPg+AoSud1I1iyJHI69jH4=
|
||||
github.com/pion/randutil v0.1.0 h1:CFG1UdESneORglEsnimhUjf33Rwjubwj6xfiOXBa3mA=
|
||||
github.com/pion/randutil v0.1.0/go.mod h1:XcJrSMMbbMRhASFVOlj/5hQial/Y8oH/HVo7TBZq+j8=
|
||||
github.com/pion/rtp v1.8.7 h1:qslKkG8qxvQ7hqaxkmL7Pl0XcUm+/Er7nMnu6Vq+ZxM=
|
||||
github.com/pion/rtp v1.8.7/go.mod h1:pBGHaFt/yW7bf1jjWAoUjpSNoDnw98KTMg+jWWvziqU=
|
||||
github.com/pion/webrtc/v3 v3.3.6 h1:7XAh4RPtlY1Vul6/GmZrv7z+NnxKA6If0KStXBI2ZLE=
|
||||
github.com/pion/webrtc/v3 v3.3.6/go.mod h1:zyN7th4mZpV27eXybfR/cnUf3J2DRy8zw/mdjD9JTNM=
|
||||
github.com/pkg/diff v0.0.0-20210226163009-20ebb0f2a09e/go.mod h1:pJLUxLENpZxwdsKMEsNbx1VGcRFpLqf3715MtcvvzbA=
|
||||
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
|
|
@ -230,6 +236,8 @@ github.com/vektah/gqlparser/v2 v2.5.27 h1:RHPD3JOplpk5mP5JGX8RKZkt2/Vwj/PZv0HxTd
|
|||
github.com/vektah/gqlparser/v2 v2.5.27/go.mod h1:D1/VCZtV3LPnQrcPBeR/q5jkSQIPti0uYCP/RI0gIeo=
|
||||
github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU=
|
||||
github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E=
|
||||
github.com/yeongaori/discordgo-fork v0.0.0-20260319072544-e8e546f5d532 h1:gxFHYeUDGziRb0zXYEqBFohC+NJbIW9L0tddaXMWr2o=
|
||||
github.com/yeongaori/discordgo-fork v0.0.0-20260319072544-e8e546f5d532/go.mod h1:A0FcMFJKJ9fRjgSuZ2o+pIQ6mPS81SVuiLN2vYTa7Ao=
|
||||
github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4=
|
||||
github.com/yosida95/uritemplate/v3 v3.0.2/go.mod h1:ILOh0sOhIJR3+L/8afwt/kE++YT040gmv5BQTMR2HP4=
|
||||
github.com/yuin/goldmark v1.1.27/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74=
|
||||
|
|
@ -249,7 +257,6 @@ golang.org/x/arch v0.24.0/go.mod h1:dNHoOeKiyja7GTvF9NJS1l3Z2yntpQNzgrjh1cU103A=
|
|||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||
golang.org/x/crypto v0.0.0-20210421170649-83a5a9bb288b/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
||||
golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc=
|
||||
golang.org/x/crypto v0.16.0/go.mod h1:gCAAfMLgwOJRpTjQ2zCCt2OcSfYMTeZVSRtQlPC7Nq4=
|
||||
golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4=
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ import (
|
|||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/asr"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/commands"
|
||||
|
|
@ -31,7 +32,6 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/state"
|
||||
"github.com/sipeed/picoclaw/pkg/tools"
|
||||
"github.com/sipeed/picoclaw/pkg/utils"
|
||||
"github.com/sipeed/picoclaw/pkg/voice"
|
||||
)
|
||||
|
||||
type AgentLoop struct {
|
||||
|
|
@ -51,7 +51,7 @@ type AgentLoop struct {
|
|||
fallback *providers.FallbackChain
|
||||
channelManager *channels.Manager
|
||||
mediaStore media.MediaStore
|
||||
transcriber voice.Transcriber
|
||||
transcriber asr.Transcriber
|
||||
cmdRegistry *commands.Registry
|
||||
mcp mcpRuntime
|
||||
hookRuntime hookRuntime
|
||||
|
|
@ -1020,7 +1020,7 @@ func (al *AgentLoop) SetMediaStore(s media.MediaStore) {
|
|||
}
|
||||
|
||||
// SetTranscriber injects a voice transcriber for agent-level audio transcription.
|
||||
func (al *AgentLoop) SetTranscriber(t voice.Transcriber) {
|
||||
func (al *AgentLoop) SetTranscriber(t asr.Transcriber) {
|
||||
al.transcriber = t
|
||||
}
|
||||
|
||||
|
|
@ -2155,147 +2155,74 @@ turnLoop:
|
|||
})
|
||||
}
|
||||
messages = append(messages, assistantMsg)
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, assistantMsg)
|
||||
ts.recordPersistedMessage(assistantMsg)
|
||||
|
||||
// Save assistant message with tool calls to session
|
||||
agent.Sessions.AddFullMessage(opts.SessionKey, assistantMsg)
|
||||
|
||||
// Execute tool calls in parallel
|
||||
type indexedAgentResult struct {
|
||||
result *tools.ToolResult
|
||||
tc providers.ToolCall
|
||||
}
|
||||
|
||||
ts.setPhase(TurnPhaseTools)
|
||||
agentResults := make([]indexedAgentResult, len(normalizedToolCalls))
|
||||
var wg sync.WaitGroup
|
||||
|
||||
for i, tc := range normalizedToolCalls {
|
||||
if ts.hardAbortRequested() {
|
||||
turnStatus = TurnEndStatusAborted
|
||||
return al.abortTurn(ts)
|
||||
}
|
||||
agentResults[i].tc = tc
|
||||
|
||||
toolName := tc.Name
|
||||
toolArgs := cloneStringAnyMap(tc.Arguments)
|
||||
wg.Add(1)
|
||||
go func(idx int, tc providers.ToolCall) {
|
||||
defer wg.Done()
|
||||
|
||||
if al.hooks != nil {
|
||||
toolReq, decision := al.hooks.BeforeTool(turnCtx, &ToolCallHookRequest{
|
||||
Meta: ts.eventMeta("runTurn", "turn.tool.before"),
|
||||
Tool: toolName,
|
||||
Arguments: toolArgs,
|
||||
Channel: ts.channel,
|
||||
ChatID: ts.chatID,
|
||||
})
|
||||
switch decision.normalizedAction() {
|
||||
case HookActionContinue, HookActionModify:
|
||||
if toolReq != nil {
|
||||
toolName = toolReq.Tool
|
||||
toolArgs = toolReq.Arguments
|
||||
}
|
||||
case HookActionDenyTool:
|
||||
denyContent := hookDeniedToolContent("Tool execution denied by hook", decision.Reason)
|
||||
al.emitEvent(
|
||||
EventKindToolExecSkipped,
|
||||
ts.eventMeta("runTurn", "turn.tool.skipped"),
|
||||
ToolExecSkippedPayload{
|
||||
Tool: toolName,
|
||||
Reason: denyContent,
|
||||
},
|
||||
)
|
||||
deniedMsg := providers.Message{
|
||||
Role: "tool",
|
||||
Content: denyContent,
|
||||
ToolCallID: tc.ID,
|
||||
}
|
||||
messages = append(messages, deniedMsg)
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, deniedMsg)
|
||||
ts.recordPersistedMessage(deniedMsg)
|
||||
}
|
||||
continue
|
||||
case HookActionAbortTurn:
|
||||
turnStatus = TurnEndStatusError
|
||||
return turnResult{}, al.hookAbortError(ts, "before_tool", decision)
|
||||
case HookActionHardAbort:
|
||||
_ = ts.requestHardAbort()
|
||||
turnStatus = TurnEndStatusAborted
|
||||
return al.abortTurn(ts)
|
||||
}
|
||||
}
|
||||
|
||||
if al.hooks != nil {
|
||||
approval := al.hooks.ApproveTool(turnCtx, &ToolApprovalRequest{
|
||||
Meta: ts.eventMeta("runTurn", "turn.tool.approve"),
|
||||
Tool: toolName,
|
||||
Arguments: toolArgs,
|
||||
Channel: ts.channel,
|
||||
ChatID: ts.chatID,
|
||||
})
|
||||
if !approval.Approved {
|
||||
denyContent := hookDeniedToolContent("Tool execution denied by approval hook", approval.Reason)
|
||||
al.emitEvent(
|
||||
EventKindToolExecSkipped,
|
||||
ts.eventMeta("runTurn", "turn.tool.skipped"),
|
||||
ToolExecSkippedPayload{
|
||||
Tool: toolName,
|
||||
Reason: denyContent,
|
||||
},
|
||||
)
|
||||
deniedMsg := providers.Message{
|
||||
Role: "tool",
|
||||
Content: denyContent,
|
||||
ToolCallID: tc.ID,
|
||||
}
|
||||
messages = append(messages, deniedMsg)
|
||||
if !ts.opts.NoHistory {
|
||||
ts.agent.Sessions.AddFullMessage(ts.sessionKey, deniedMsg)
|
||||
ts.recordPersistedMessage(deniedMsg)
|
||||
}
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
argsJSON, _ := json.Marshal(toolArgs)
|
||||
argsPreview := utils.Truncate(string(argsJSON), 200)
|
||||
logger.InfoCF("agent", fmt.Sprintf("Tool call: %s(%s)", toolName, argsPreview),
|
||||
map[string]any{
|
||||
"agent_id": ts.agent.ID,
|
||||
"tool": toolName,
|
||||
"iteration": iteration,
|
||||
})
|
||||
al.emitEvent(
|
||||
EventKindToolExecStart,
|
||||
ts.eventMeta("runTurn", "turn.tool.start"),
|
||||
ToolExecStartPayload{
|
||||
Tool: toolName,
|
||||
Arguments: cloneEventArguments(toolArgs),
|
||||
},
|
||||
)
|
||||
|
||||
// Send tool feedback to chat channel if enabled (from HEAD)
|
||||
if al.cfg.Agents.Defaults.IsToolFeedbackEnabled() && ts.channel != "" {
|
||||
feedbackPreview := utils.Truncate(
|
||||
string(argsJSON),
|
||||
al.cfg.Agents.Defaults.GetToolFeedbackMaxArgsLength(),
|
||||
)
|
||||
feedbackMsg := fmt.Sprintf("\U0001f527 `%s`\n```\n%s\n```", tc.Name, feedbackPreview)
|
||||
fbCtx, fbCancel := context.WithTimeout(turnCtx, 3*time.Second)
|
||||
_ = al.bus.PublishOutbound(fbCtx, bus.OutboundMessage{
|
||||
Channel: ts.channel,
|
||||
ChatID: ts.chatID,
|
||||
Content: feedbackMsg,
|
||||
})
|
||||
fbCancel()
|
||||
}
|
||||
|
||||
toolCallID := tc.ID
|
||||
toolIteration := iteration
|
||||
asyncToolName := toolName
|
||||
asyncCallback := func(_ context.Context, result *tools.ToolResult) {
|
||||
// Send ForUser content directly to the user (immediate feedback),
|
||||
// mirroring the synchronous tool execution path.
|
||||
if !result.Silent && result.ForUser != "" {
|
||||
outCtx, outCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer outCancel()
|
||||
_ = al.bus.PublishOutbound(outCtx, bus.OutboundMessage{
|
||||
Channel: ts.channel,
|
||||
ChatID: ts.chatID,
|
||||
Content: result.ForUser,
|
||||
argsJSON, _ := json.Marshal(tc.Arguments)
|
||||
argsPreview := utils.Truncate(string(argsJSON), 200)
|
||||
logger.InfoCF("agent", fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview),
|
||||
map[string]any{
|
||||
"agent_id": agent.ID,
|
||||
"tool": tc.Name,
|
||||
"iteration": iteration,
|
||||
})
|
||||
|
||||
// Send tool feedback to chat channel if enabled
|
||||
if al.cfg.Agents.Defaults.IsToolFeedbackEnabled() && opts.Channel != "" {
|
||||
feedbackPreview := utils.Truncate(
|
||||
string(argsJSON),
|
||||
al.cfg.Agents.Defaults.GetToolFeedbackMaxArgsLength(),
|
||||
)
|
||||
feedbackMsg := fmt.Sprintf("\U0001f527 `%s`\n```\n%s\n```", tc.Name, feedbackPreview)
|
||||
fbCtx, fbCancel := context.WithTimeout(ctx, 3*time.Second)
|
||||
_ = al.bus.PublishOutbound(fbCtx, bus.OutboundMessage{
|
||||
Channel: opts.Channel,
|
||||
ChatID: opts.ChatID,
|
||||
Content: feedbackMsg,
|
||||
Metadata: map[string]string{
|
||||
"is_tool_call": "true",
|
||||
},
|
||||
})
|
||||
fbCancel()
|
||||
}
|
||||
|
||||
// Create async callback for tools that implement AsyncExecutor.
|
||||
// When the background work completes, this publishes the result
|
||||
// as an inbound system message so processSystemMessage routes it
|
||||
// back to the user via the normal agent loop.
|
||||
asyncCallback := func(_ context.Context, result *tools.ToolResult) {
|
||||
// Send ForUser content directly to the user (immediate feedback),
|
||||
// mirroring the synchronous tool execution path.
|
||||
if !result.Silent && result.ForUser != "" {
|
||||
outCtx, outCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer outCancel()
|
||||
_ = al.bus.PublishOutbound(outCtx, bus.OutboundMessage{
|
||||
Channel: opts.Channel,
|
||||
ChatID: opts.ChatID,
|
||||
Content: result.ForUser,
|
||||
Metadata: map[string]string{
|
||||
"is_tool_call": "true",
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Determine content for the agent loop (ForLLM or error).
|
||||
content := result.ForLLM
|
||||
if content == "" && result.Err != nil {
|
||||
|
|
@ -2384,9 +2311,12 @@ turnLoop:
|
|||
|
||||
if !toolResult.Silent && toolResult.ForUser != "" && ts.opts.SendResponse {
|
||||
al.bus.PublishOutbound(ctx, bus.OutboundMessage{
|
||||
Channel: ts.channel,
|
||||
ChatID: ts.chatID,
|
||||
Content: toolResult.ForUser,
|
||||
Channel: opts.Channel,
|
||||
ChatID: opts.ChatID,
|
||||
Content: r.result.ForUser,
|
||||
Metadata: map[string]string{
|
||||
"is_tool_call": "true",
|
||||
},
|
||||
})
|
||||
logger.DebugCF("agent", "Sent tool result to user",
|
||||
map[string]any{
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
package voice
|
||||
package asr
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
|
|
@ -21,6 +21,7 @@ import (
|
|||
type Transcriber interface {
|
||||
Name() string
|
||||
Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error)
|
||||
TranscribeData(ctx context.Context, data []byte, filename string) (*TranscriptionResponse, error)
|
||||
}
|
||||
|
||||
type GroqTranscriber struct {
|
||||
|
|
@ -48,45 +49,28 @@ func NewGroqTranscriber(apiKey string) *GroqTranscriber {
|
|||
}
|
||||
}
|
||||
|
||||
func (t *GroqTranscriber) Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error) {
|
||||
logger.InfoCF("voice", "Starting transcription", map[string]any{"audio_file": audioFilePath})
|
||||
|
||||
audioFile, err := os.Open(audioFilePath)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to open audio file", map[string]any{"path": audioFilePath, "error": err})
|
||||
return nil, fmt.Errorf("failed to open audio file: %w", err)
|
||||
}
|
||||
defer audioFile.Close()
|
||||
|
||||
fileInfo, err := audioFile.Stat()
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to get file info", map[string]any{"path": audioFilePath, "error": err})
|
||||
return nil, fmt.Errorf("failed to get file info: %w", err)
|
||||
}
|
||||
|
||||
logger.DebugCF("voice", "Audio file details", map[string]any{
|
||||
"size_bytes": fileInfo.Size(),
|
||||
"file_name": filepath.Base(audioFilePath),
|
||||
})
|
||||
func (t *GroqTranscriber) TranscribeData(
|
||||
ctx context.Context,
|
||||
data []byte,
|
||||
filename string,
|
||||
) (*TranscriptionResponse, error) {
|
||||
logger.InfoCF("voice", "Starting memory transcription", map[string]any{"filename": filename, "bytes": len(data)})
|
||||
|
||||
var requestBody bytes.Buffer
|
||||
writer := multipart.NewWriter(&requestBody)
|
||||
|
||||
part, err := writer.CreateFormFile("file", filepath.Base(audioFilePath))
|
||||
part, err := writer.CreateFormFile("file", filename)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to create form file", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to create form file: %w", err)
|
||||
}
|
||||
|
||||
copied, err := io.Copy(part, audioFile)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to copy file content", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to copy file content: %w", err)
|
||||
if _, copyErr := io.Copy(part, bytes.NewReader(data)); copyErr != nil {
|
||||
logger.ErrorCF("voice", "Failed to copy file content", map[string]any{"error": copyErr})
|
||||
return nil, fmt.Errorf("failed to copy file content: %w", copyErr)
|
||||
}
|
||||
|
||||
logger.DebugCF("voice", "File copied to request", map[string]any{"bytes_copied": copied})
|
||||
|
||||
if err = writer.WriteField("model", "whisper-large-v3"); err != nil {
|
||||
if err = writer.WriteField("model", "whisper-large-v3-turbo"); err != nil {
|
||||
logger.ErrorCF("voice", "Failed to write model field", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to write model field: %w", err)
|
||||
}
|
||||
|
|
@ -101,20 +85,70 @@ func (t *GroqTranscriber) Transcribe(ctx context.Context, audioFilePath string)
|
|||
return nil, fmt.Errorf("failed to close multipart writer: %w", err)
|
||||
}
|
||||
|
||||
return t.doRequest(ctx, &requestBody, writer.FormDataContentType(), int64(len(data)))
|
||||
}
|
||||
|
||||
func (t *GroqTranscriber) Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error) {
|
||||
logger.InfoCF("voice", "Starting transcription", map[string]any{"audio_file": audioFilePath})
|
||||
|
||||
audioFile, err := os.Open(audioFilePath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to open audio file %s: %w", audioFilePath, err)
|
||||
}
|
||||
defer audioFile.Close()
|
||||
|
||||
fileInfo, err := audioFile.Stat()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to stat audio file %s: %w", audioFilePath, err)
|
||||
}
|
||||
|
||||
var requestBody bytes.Buffer
|
||||
writer := multipart.NewWriter(&requestBody)
|
||||
|
||||
part, err := writer.CreateFormFile("file", filepath.Base(audioFilePath))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create form file: %w", err)
|
||||
}
|
||||
|
||||
if _, copyErr := io.Copy(part, audioFile); copyErr != nil {
|
||||
return nil, fmt.Errorf("failed to copy audio data: %w", copyErr)
|
||||
}
|
||||
|
||||
if err = writer.WriteField("model", "whisper-large-v3-turbo"); err != nil {
|
||||
return nil, fmt.Errorf("failed to write model field: %w", err)
|
||||
}
|
||||
|
||||
if err = writer.WriteField("response_format", "json"); err != nil {
|
||||
return nil, fmt.Errorf("failed to write response_format field: %w", err)
|
||||
}
|
||||
|
||||
if err = writer.Close(); err != nil {
|
||||
return nil, fmt.Errorf("failed to close multipart writer: %w", err)
|
||||
}
|
||||
|
||||
return t.doRequest(ctx, &requestBody, writer.FormDataContentType(), fileInfo.Size())
|
||||
}
|
||||
|
||||
func (t *GroqTranscriber) doRequest(
|
||||
ctx context.Context,
|
||||
requestBody *bytes.Buffer,
|
||||
contentType string,
|
||||
fileSize int64,
|
||||
) (*TranscriptionResponse, error) {
|
||||
url := t.apiBase + "/audio/transcriptions"
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", url, &requestBody)
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", url, requestBody)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice", "Failed to create request", map[string]any{"error": err})
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
req.Header.Set("Content-Type", contentType)
|
||||
req.Header.Set("Authorization", "Bearer "+t.apiKey)
|
||||
|
||||
logger.DebugCF("voice", "Sending transcription request to Groq API", map[string]any{
|
||||
"url": url,
|
||||
"request_size_bytes": requestBody.Len(),
|
||||
"file_size_bytes": fileInfo.Size(),
|
||||
"file_size_bytes": fileSize,
|
||||
})
|
||||
|
||||
resp, err := t.httpClient.Do(req)
|
||||
|
|
@ -170,9 +204,10 @@ func DetectTranscriber(cfg *config.Config) Transcriber {
|
|||
if key := cfg.Providers.Groq.APIKey; key != "" {
|
||||
return NewGroqTranscriber(key)
|
||||
}
|
||||
// Fall back to any model-list entry that uses the groq/ protocol.
|
||||
// Fall back to any model-list entry that uses the groq/ protocol or is explicitly named groq.
|
||||
for _, mc := range cfg.ModelList {
|
||||
if strings.HasPrefix(mc.Model, "groq/") && mc.APIKey != "" {
|
||||
if (strings.HasPrefix(mc.Model, "groq/") || mc.ModelName == "groq" || mc.Model == "whisper-large-v3-turbo") &&
|
||||
mc.APIKey != "" {
|
||||
return NewGroqTranscriber(mc.APIKey)
|
||||
}
|
||||
}
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
package voice
|
||||
package asr
|
||||
|
||||
import (
|
||||
"context"
|
||||
55
pkg/audio/ogg.go
Normal file
55
pkg/audio/ogg.go
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
package audio
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
// DecodeOggOpus reads an Ogg format stream and extracts individual Opus payloads.
|
||||
// It calls onFrame for every complete Opus frame found in the stream.
|
||||
func DecodeOggOpus(r io.Reader, onFrame func([]byte) error) error {
|
||||
var packet []byte
|
||||
header := make([]byte, 27)
|
||||
|
||||
for {
|
||||
if _, err := io.ReadFull(r, header); err != nil {
|
||||
if err == io.EOF || err == io.ErrUnexpectedEOF {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("failed to read ogg header: %w", err)
|
||||
}
|
||||
if string(header[:4]) != "OggS" {
|
||||
return fmt.Errorf("invalid ogg magic string")
|
||||
}
|
||||
|
||||
pageSegments := int(header[26])
|
||||
segmentTable := make([]byte, pageSegments)
|
||||
if _, err := io.ReadFull(r, segmentTable); err != nil {
|
||||
return fmt.Errorf("failed to read segment table: %w", err)
|
||||
}
|
||||
|
||||
for _, lacing := range segmentTable {
|
||||
segment := make([]byte, lacing)
|
||||
if _, err := io.ReadFull(r, segment); err != nil {
|
||||
return fmt.Errorf("failed to read segment data: %w", err)
|
||||
}
|
||||
|
||||
packet = append(packet, segment...)
|
||||
|
||||
// If lacing is less than 255, the packet is complete
|
||||
if lacing < 255 {
|
||||
if len(packet) > 0 {
|
||||
// Ignore Ogg Opus headers
|
||||
if !bytes.HasPrefix(packet, []byte("OpusHead")) && !bytes.HasPrefix(packet, []byte("OpusTags")) {
|
||||
if err := onFrame(packet); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
// Start new packet
|
||||
packet = nil
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
96
pkg/audio/sentence.go
Normal file
96
pkg/audio/sentence.go
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
package audio
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
// SplitSentences splits text into sentence-sized chunks suitable for TTS synthesis.
|
||||
// It splits on sentence-ending punctuation (.!?\n) while avoiding false splits
|
||||
// on abbreviations and decimal numbers. Very short fragments are merged with
|
||||
// the next sentence to prevent choppy playback.
|
||||
func SplitSentences(text string) []string {
|
||||
if text == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
var sentences []string
|
||||
var current strings.Builder
|
||||
runes := []rune(text)
|
||||
|
||||
for i := 0; i < len(runes); i++ {
|
||||
r := runes[i]
|
||||
current.WriteRune(r)
|
||||
|
||||
if r == '\n' {
|
||||
s := strings.TrimSpace(current.String())
|
||||
if s != "" {
|
||||
sentences = append(sentences, s)
|
||||
}
|
||||
current.Reset()
|
||||
continue
|
||||
}
|
||||
|
||||
if r == '.' || r == '!' || r == '?' {
|
||||
// Avoid splitting on decimal numbers like "3.14"
|
||||
if r == '.' && i > 0 && unicode.IsDigit(runes[i-1]) &&
|
||||
i+1 < len(runes) && unicode.IsDigit(runes[i+1]) {
|
||||
continue
|
||||
}
|
||||
|
||||
// Consume trailing punctuation and spaces (e.g., "..." or "?!")
|
||||
for i+1 < len(runes) && (runes[i+1] == '.' || runes[i+1] == '!' || runes[i+1] == '?' || runes[i+1] == ' ') {
|
||||
i++
|
||||
current.WriteRune(runes[i])
|
||||
}
|
||||
|
||||
s := strings.TrimSpace(current.String())
|
||||
if s != "" {
|
||||
sentences = append(sentences, s)
|
||||
}
|
||||
current.Reset()
|
||||
}
|
||||
}
|
||||
|
||||
// Flush remaining text
|
||||
if s := strings.TrimSpace(current.String()); s != "" {
|
||||
sentences = append(sentences, s)
|
||||
}
|
||||
|
||||
// Merge very short fragments with the next sentence
|
||||
return mergeShorties(sentences, 15)
|
||||
}
|
||||
|
||||
// mergeShorties merges sentences shorter than minLen characters with the following sentence.
|
||||
func mergeShorties(sentences []string, minLen int) []string {
|
||||
if len(sentences) <= 1 {
|
||||
return sentences
|
||||
}
|
||||
|
||||
var merged []string
|
||||
var buf string
|
||||
|
||||
for _, s := range sentences {
|
||||
if buf != "" {
|
||||
buf += " " + s
|
||||
if len([]rune(buf)) >= minLen {
|
||||
merged = append(merged, buf)
|
||||
buf = ""
|
||||
}
|
||||
} else if len([]rune(s)) < minLen {
|
||||
buf = s
|
||||
} else {
|
||||
merged = append(merged, s)
|
||||
}
|
||||
}
|
||||
|
||||
if buf != "" {
|
||||
if len(merged) > 0 {
|
||||
merged[len(merged)-1] += " " + buf
|
||||
} else {
|
||||
merged = append(merged, buf)
|
||||
}
|
||||
}
|
||||
|
||||
return merged
|
||||
}
|
||||
|
|
@ -34,6 +34,8 @@ type MessageBus struct {
|
|||
inbound chan InboundMessage
|
||||
outbound chan OutboundMessage
|
||||
outboundMedia chan OutboundMediaMessage
|
||||
audioChunks chan AudioChunk
|
||||
voiceControls chan VoiceControl
|
||||
|
||||
closeOnce sync.Once
|
||||
done chan struct{}
|
||||
|
|
@ -47,6 +49,8 @@ func NewMessageBus() *MessageBus {
|
|||
inbound: make(chan InboundMessage, defaultBusBufferSize),
|
||||
outbound: make(chan OutboundMessage, defaultBusBufferSize),
|
||||
outboundMedia: make(chan OutboundMediaMessage, defaultBusBufferSize),
|
||||
audioChunks: make(chan AudioChunk, defaultBusBufferSize*4), // Audio chunks need more buffer
|
||||
voiceControls: make(chan VoiceControl, defaultBusBufferSize),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
|
@ -103,6 +107,22 @@ func (mb *MessageBus) OutboundMediaChan() <-chan OutboundMediaMessage {
|
|||
return mb.outboundMedia
|
||||
}
|
||||
|
||||
func (mb *MessageBus) PublishAudioChunk(ctx context.Context, chunk AudioChunk) error {
|
||||
return publish(ctx, mb, mb.audioChunks, chunk)
|
||||
}
|
||||
|
||||
func (mb *MessageBus) AudioChunksChan() <-chan AudioChunk {
|
||||
return mb.audioChunks
|
||||
}
|
||||
|
||||
func (mb *MessageBus) PublishVoiceControl(ctx context.Context, ctrl VoiceControl) error {
|
||||
return publish(ctx, mb, mb.voiceControls, ctrl)
|
||||
}
|
||||
|
||||
func (mb *MessageBus) VoiceControlsChan() <-chan VoiceControl {
|
||||
return mb.voiceControls
|
||||
}
|
||||
|
||||
// SetStreamDelegate registers a StreamDelegate (typically the channel Manager).
|
||||
func (mb *MessageBus) SetStreamDelegate(d StreamDelegate) {
|
||||
mb.streamDelegate.Store(d)
|
||||
|
|
@ -132,6 +152,8 @@ func (mb *MessageBus) Close() {
|
|||
close(mb.inbound)
|
||||
close(mb.outbound)
|
||||
close(mb.outboundMedia)
|
||||
close(mb.audioChunks)
|
||||
close(mb.voiceControls)
|
||||
|
||||
// clean up any remaining messages in channels
|
||||
drained := 0
|
||||
|
|
@ -144,6 +166,12 @@ func (mb *MessageBus) Close() {
|
|||
for range mb.outboundMedia {
|
||||
drained++
|
||||
}
|
||||
for range mb.audioChunks {
|
||||
drained++
|
||||
}
|
||||
for range mb.voiceControls {
|
||||
drained++
|
||||
}
|
||||
|
||||
if drained > 0 {
|
||||
logger.DebugCF("bus", "Drained buffered messages during close", map[string]any{
|
||||
|
|
|
|||
|
|
@ -30,10 +30,11 @@ type InboundMessage struct {
|
|||
}
|
||||
|
||||
type OutboundMessage struct {
|
||||
Channel string `json:"channel"`
|
||||
ChatID string `json:"chat_id"`
|
||||
Content string `json:"content"`
|
||||
ReplyToMessageID string `json:"reply_to_message_id,omitempty"`
|
||||
Channel string `json:"channel"`
|
||||
ChatID string `json:"chat_id"`
|
||||
Content string `json:"content"`
|
||||
ReplyToMessageID string `json:"reply_to_message_id,omitempty"`
|
||||
Metadata map[string]string `json:"metadata,omitempty"`
|
||||
}
|
||||
|
||||
// MediaPart describes a single media attachment to send.
|
||||
|
|
@ -51,3 +52,24 @@ type OutboundMediaMessage struct {
|
|||
ChatID string `json:"chat_id"`
|
||||
Parts []MediaPart `json:"parts"`
|
||||
}
|
||||
|
||||
// AudioChunk represents a chunk of streaming voice data.
|
||||
type AudioChunk struct {
|
||||
SessionID string `json:"session_id"`
|
||||
SpeakerID string `json:"speaker_id"` // User ID or SSRC
|
||||
ChatID string `json:"chat_id"` // Where to respond
|
||||
Channel string `json:"channel"` // Source channel type (e.g. "discord")
|
||||
Sequence uint64 `json:"sequence"`
|
||||
Timestamp uint32 `json:"timestamp"`
|
||||
SampleRate int `json:"sample_rate"`
|
||||
Channels int `json:"channels"`
|
||||
Format string `json:"format"` // "opus", "pcm", etc
|
||||
Data []byte `json:"data"`
|
||||
}
|
||||
|
||||
// VoiceControl represents state or commands for voice sessions.
|
||||
type VoiceControl struct {
|
||||
SessionID string `json:"session_id"`
|
||||
Type string `json:"type"` // "state", "command"
|
||||
Action string `json:"action"` // "idle", "listening", "start", "stop", "leave"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ package discord
|
|||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
|
|
@ -14,12 +15,14 @@ import (
|
|||
"github.com/bwmarrin/discordgo"
|
||||
"github.com/gorilla/websocket"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/audio"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/identity"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
"github.com/sipeed/picoclaw/pkg/media"
|
||||
"github.com/sipeed/picoclaw/pkg/tts"
|
||||
"github.com/sipeed/picoclaw/pkg/utils"
|
||||
)
|
||||
|
||||
|
|
@ -42,6 +45,12 @@ type DiscordChannel struct {
|
|||
typingMu sync.Mutex
|
||||
typingStop map[string]chan struct{} // chatID → stop signal
|
||||
botUserID string // stored for mention checking
|
||||
bus *bus.MessageBus
|
||||
tts tts.TTSProvider
|
||||
|
||||
// TTS interruption: cancel active playback when user speaks
|
||||
ttsMu sync.Mutex
|
||||
cancelTTS context.CancelFunc
|
||||
}
|
||||
|
||||
func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordChannel, error) {
|
||||
|
|
@ -73,6 +82,7 @@ func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordC
|
|||
config: cfg,
|
||||
ctx: context.Background(),
|
||||
typingStop: make(map[string]chan struct{}),
|
||||
bus: bus,
|
||||
}, nil
|
||||
}
|
||||
|
||||
|
|
@ -90,6 +100,8 @@ func (c *DiscordChannel) Start(ctx context.Context) error {
|
|||
|
||||
c.session.AddHandler(c.handleMessage)
|
||||
|
||||
go c.listenVoiceControl(c.ctx)
|
||||
|
||||
if err := c.session.Open(); err != nil {
|
||||
return fmt.Errorf("failed to open discord session: %w", err)
|
||||
}
|
||||
|
|
@ -142,6 +154,30 @@ func (c *DiscordChannel) Send(ctx context.Context, msg bus.OutboundMessage) erro
|
|||
return nil
|
||||
}
|
||||
|
||||
isToolCall := false
|
||||
if msg.Metadata != nil {
|
||||
if val, ok := msg.Metadata["is_tool_call"]; ok && val == "true" {
|
||||
isToolCall = true
|
||||
}
|
||||
}
|
||||
|
||||
if c.tts != nil && !isToolCall {
|
||||
if ch, err := c.session.State.Channel(channelID); err == nil && ch.GuildID != "" {
|
||||
if vc, ok := c.session.VoiceConnections[ch.GuildID]; ok && vc != nil {
|
||||
// Cancel any previous TTS playback
|
||||
c.ttsMu.Lock()
|
||||
if c.cancelTTS != nil {
|
||||
c.cancelTTS()
|
||||
}
|
||||
ttsCtx, ttsCancel := context.WithCancel(c.ctx)
|
||||
c.cancelTTS = ttsCancel
|
||||
c.ttsMu.Unlock()
|
||||
|
||||
go c.playTTS(ttsCtx, vc, msg.Content)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return c.sendChunk(ctx, channelID, msg.Content, msg.ReplyToMessageID)
|
||||
}
|
||||
|
||||
|
|
@ -321,6 +357,10 @@ func (c *DiscordChannel) handleMessage(s *discordgo.Session, m *discordgo.Messag
|
|||
return
|
||||
}
|
||||
|
||||
if c.handleVoiceCommand(s, m) {
|
||||
return
|
||||
}
|
||||
|
||||
// Check allowlist first to avoid downloading attachments for rejected users
|
||||
sender := bus.SenderInfo{
|
||||
Platform: "discord",
|
||||
|
|
@ -612,3 +652,106 @@ func (c *DiscordChannel) stripBotMention(text string) string {
|
|||
text = strings.ReplaceAll(text, fmt.Sprintf("<@!%s>", c.botUserID), "")
|
||||
return strings.TrimSpace(text)
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) listenVoiceControl(ctx context.Context) {
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ctrl, ok := <-c.bus.VoiceControlsChan():
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if ctrl.Type == "command" && ctrl.Action == "leave" {
|
||||
if strings.HasPrefix(ctrl.SessionID, "discord_vc_") {
|
||||
guildID := strings.TrimPrefix(ctrl.SessionID, "discord_vc_")
|
||||
vc, exists := c.session.VoiceConnections[guildID]
|
||||
if exists && vc != nil {
|
||||
vc.Disconnect(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) playTTS(ctx context.Context, vc *discordgo.VoiceConnection, text string) {
|
||||
// Clear cancelTTS when playback finishes (normal or interrupted)
|
||||
defer func() {
|
||||
c.ttsMu.Lock()
|
||||
c.cancelTTS = nil
|
||||
c.ttsMu.Unlock()
|
||||
}()
|
||||
|
||||
sentences := audio.SplitSentences(text)
|
||||
if len(sentences) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
logger.InfoCF("discord", "Starting streamed TTS", map[string]any{"sentences": len(sentences)})
|
||||
|
||||
// Pipeline: prefetch next sentence's audio while playing current
|
||||
type ttResult struct {
|
||||
stream io.ReadCloser
|
||||
err error
|
||||
}
|
||||
|
||||
var prefetch chan ttResult
|
||||
|
||||
// Ensure any in-flight prefetch is drained on exit to prevent stream leaks
|
||||
defer func() {
|
||||
if prefetch != nil {
|
||||
result := <-prefetch
|
||||
if result.stream != nil {
|
||||
result.stream.Close()
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
for i, sentence := range sentences {
|
||||
// Check for cancellation (interruption)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
logger.InfoCF("discord", "TTS interrupted", map[string]any{"at_sentence": i})
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
// Start prefetching the NEXT sentence while we process the current one
|
||||
var nextPrefetch chan ttResult
|
||||
if i+1 < len(sentences) {
|
||||
nextPrefetch = make(chan ttResult, 1)
|
||||
nextSentence := sentences[i+1]
|
||||
go func() {
|
||||
s, e := c.tts.Synthesize(ctx, nextSentence)
|
||||
nextPrefetch <- ttResult{s, e}
|
||||
}()
|
||||
}
|
||||
|
||||
// Get the current sentence's audio
|
||||
var stream io.ReadCloser
|
||||
var err error
|
||||
|
||||
if prefetch != nil {
|
||||
// Use prefetched result from previous iteration
|
||||
result := <-prefetch
|
||||
stream, err = result.stream, result.err
|
||||
} else {
|
||||
// First sentence: synthesize directly
|
||||
stream, err = c.tts.Synthesize(ctx, sentence)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
logger.ErrorCF("discord", "TTS synthesize failed", map[string]any{"error": err.Error(), "sentence": i})
|
||||
prefetch = nextPrefetch
|
||||
continue
|
||||
}
|
||||
|
||||
if err := streamOggOpusToDiscord(ctx, vc, stream); err != nil {
|
||||
logger.ErrorCF("discord", "TTS playback failed", map[string]any{"error": err.Error(), "sentence": i})
|
||||
}
|
||||
stream.Close()
|
||||
|
||||
prefetch = nextPrefetch
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,10 +4,15 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/tts"
|
||||
)
|
||||
|
||||
func init() {
|
||||
channels.RegisterFactory("discord", func(cfg *config.Config, b *bus.MessageBus) (channels.Channel, error) {
|
||||
return NewDiscordChannel(cfg.Channels.Discord, b)
|
||||
ch, err := NewDiscordChannel(cfg.Channels.Discord, b)
|
||||
if err == nil {
|
||||
ch.tts = tts.DetectTTS(cfg)
|
||||
}
|
||||
return ch, err
|
||||
})
|
||||
}
|
||||
|
|
|
|||
211
pkg/channels/discord/voice.go
Normal file
211
pkg/channels/discord/voice.go
Normal file
|
|
@ -0,0 +1,211 @@
|
|||
package discord
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"time"
|
||||
|
||||
"github.com/bwmarrin/discordgo"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/audio"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
)
|
||||
|
||||
func (c *DiscordChannel) handleVoiceCommand(s *discordgo.Session, m *discordgo.MessageCreate) bool {
|
||||
if m.Content == "!vc join" {
|
||||
vs, err := s.State.VoiceState(m.GuildID, m.Author.ID)
|
||||
if err != nil || vs == nil {
|
||||
s.ChannelMessageSend(m.ChannelID, "You need to be in a voice channel first!")
|
||||
return true
|
||||
}
|
||||
|
||||
logger.InfoCF("discord", "Joining voice channel", map[string]any{"channel": vs.ChannelID})
|
||||
vc, err := s.ChannelVoiceJoin(c.ctx, m.GuildID, vs.ChannelID, false, false)
|
||||
if err != nil {
|
||||
s.ChannelMessageSend(m.ChannelID, fmt.Sprintf("Failed to join voice channel: %v", err))
|
||||
return true
|
||||
}
|
||||
|
||||
go c.receiveVoice(vc, m.GuildID, m.ChannelID)
|
||||
s.ChannelMessageSend(m.ChannelID, "Joined Voice Channel! Listening for audio...")
|
||||
return true
|
||||
} else if m.Content == "!vc leave" {
|
||||
vc, exists := s.VoiceConnections[m.GuildID]
|
||||
if exists && vc != nil {
|
||||
vc.Disconnect(c.ctx)
|
||||
s.ChannelMessageSend(m.ChannelID, "Left Voice Channel.")
|
||||
} else {
|
||||
s.ChannelMessageSend(m.ChannelID, "Not in a voice channel.")
|
||||
}
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func VoiceReceiveActive(vc *discordgo.VoiceConnection) bool {
|
||||
return vc != nil && vc.OpusRecv != nil
|
||||
}
|
||||
|
||||
func streamOggOpusToDiscord(ctx context.Context, vc *discordgo.VoiceConnection, r io.Reader) (retErr error) {
|
||||
// Recover from panic if vc.OpusSend is closed mid-send (e.g. on disconnect)
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
retErr = fmt.Errorf("voice connection closed during playback")
|
||||
}
|
||||
}()
|
||||
|
||||
// Wait for the speaking transition to register
|
||||
vc.Speaking(true)
|
||||
defer vc.Speaking(false)
|
||||
|
||||
return audio.DecodeOggOpus(r, func(frame []byte) error {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case vc.OpusSend <- frame:
|
||||
return nil
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func (c *DiscordChannel) receiveVoice(vc *discordgo.VoiceConnection, guildID string, chatID string) {
|
||||
logger.InfoCF("discord", "Started listening for voice", map[string]any{"guild": guildID})
|
||||
|
||||
go func(ctx context.Context, vc *discordgo.VoiceConnection) {
|
||||
// Recover from potential panics if OpusSend is closed mid-send.
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
logger.WarnCF("discord", "Recovered from panic while sending wake-up frames", map[string]any{
|
||||
"error": rec,
|
||||
"guild": guildID,
|
||||
})
|
||||
}
|
||||
}()
|
||||
|
||||
// If the voice connection or OpusSend are not available, nothing to do.
|
||||
if vc == nil || vc.OpusSend == nil {
|
||||
return
|
||||
}
|
||||
|
||||
time.Sleep(250 * time.Millisecond) // Wait a bit for connection to settle
|
||||
|
||||
// Abort if the context has already been canceled.
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
vc.Speaking(true)
|
||||
defer vc.Speaking(false)
|
||||
|
||||
silenceFrame := []byte{0xF8, 0xFF, 0xFE}
|
||||
for i := 0; i < 5; i++ {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case vc.OpusSend <- silenceFrame:
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
|
||||
logger.DebugCF("discord", "Sent wake-up silence frames", map[string]any{"guild": guildID})
|
||||
}(c.ctx, vc)
|
||||
sessionID := fmt.Sprintf("discord_vc_%s", guildID)
|
||||
|
||||
c.bus.PublishVoiceControl(c.ctx, bus.VoiceControl{
|
||||
SessionID: sessionID,
|
||||
Type: "state",
|
||||
Action: "listening",
|
||||
})
|
||||
|
||||
var sequence uint64 = 0
|
||||
var interruptCount int
|
||||
var lastInterruptAt time.Time
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-c.ctx.Done():
|
||||
return
|
||||
case p, ok := <-vc.OpusRecv:
|
||||
if !ok {
|
||||
logger.InfoCF("discord", "Voice channel closed", map[string]any{"guild": guildID})
|
||||
// Cancel any TTS that may still be playing
|
||||
c.ttsMu.Lock()
|
||||
if c.cancelTTS != nil {
|
||||
c.cancelTTS()
|
||||
c.cancelTTS = nil
|
||||
}
|
||||
c.ttsMu.Unlock()
|
||||
return
|
||||
}
|
||||
|
||||
if p == nil {
|
||||
logger.DebugCF("discord", "Received nil Opus packet", nil)
|
||||
continue
|
||||
}
|
||||
|
||||
if len(p.Opus) == 0 {
|
||||
logger.DebugCF("discord", "Received empty Opus packet", map[string]any{
|
||||
"seq": p.Sequence,
|
||||
"ssrc": p.SSRC,
|
||||
})
|
||||
continue
|
||||
}
|
||||
|
||||
logger.DebugCF("discord", "Received Opus packet", map[string]any{
|
||||
"seq": p.Sequence,
|
||||
"len": len(p.Opus),
|
||||
"ssrc": p.SSRC,
|
||||
})
|
||||
// Interruption detection: if user sends voice while TTS is playing,
|
||||
// cancel TTS after a short debounce (3 packets in 200ms)
|
||||
now := time.Now()
|
||||
if now.Sub(lastInterruptAt) > 500*time.Millisecond {
|
||||
interruptCount = 0
|
||||
}
|
||||
interruptCount++
|
||||
lastInterruptAt = now
|
||||
|
||||
if interruptCount >= 3 {
|
||||
c.ttsMu.Lock()
|
||||
if c.cancelTTS != nil {
|
||||
c.cancelTTS()
|
||||
c.cancelTTS = nil
|
||||
logger.InfoCF("discord", "TTS interrupted by user voice", nil)
|
||||
}
|
||||
c.ttsMu.Unlock()
|
||||
interruptCount = 0
|
||||
}
|
||||
|
||||
sequence++
|
||||
|
||||
chunk := bus.AudioChunk{
|
||||
SessionID: sessionID,
|
||||
SpeakerID: fmt.Sprintf("%d", p.SSRC),
|
||||
ChatID: chatID,
|
||||
Channel: "discord",
|
||||
Sequence: sequence,
|
||||
Timestamp: p.Timestamp,
|
||||
SampleRate: 48000,
|
||||
Channels: 2,
|
||||
Format: "opus",
|
||||
Data: p.Opus,
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(c.ctx, 100*time.Millisecond)
|
||||
err := c.bus.PublishAudioChunk(ctx, chunk)
|
||||
cancel()
|
||||
if err != nil {
|
||||
logger.ErrorCF("discord", "Failed to publish audio chunk", map[string]any{
|
||||
"guild": guildID,
|
||||
"sessionID": sessionID,
|
||||
"sequence": sequence,
|
||||
"error": err.Error(),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -12,6 +12,7 @@ import (
|
|||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/agent"
|
||||
"github.com/sipeed/picoclaw/pkg/asr"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/channels"
|
||||
_ "github.com/sipeed/picoclaw/pkg/channels/dingtalk"
|
||||
|
|
@ -56,6 +57,7 @@ type services struct {
|
|||
ChannelManager *channels.Manager
|
||||
DeviceService *devices.Service
|
||||
HealthServer *health.Server
|
||||
VoiceAgentCancel context.CancelFunc
|
||||
manualReloadChan chan struct{}
|
||||
reloading atomic.Bool
|
||||
}
|
||||
|
|
@ -286,9 +288,15 @@ func setupAndStartServices(
|
|||
agentLoop.SetChannelManager(runningServices.ChannelManager)
|
||||
agentLoop.SetMediaStore(runningServices.MediaStore)
|
||||
|
||||
if transcriber := voice.DetectTranscriber(cfg); transcriber != nil {
|
||||
if transcriber := asr.DetectTranscriber(cfg); transcriber != nil {
|
||||
agentLoop.SetTranscriber(transcriber)
|
||||
logger.InfoCF("voice", "Transcription enabled (agent-level)", map[string]any{"provider": transcriber.Name()})
|
||||
|
||||
// Start Voice Agent Orchestrator
|
||||
vaCtx, vaCancel := context.WithCancel(context.Background())
|
||||
runningServices.VoiceAgentCancel = vaCancel
|
||||
voiceAgent := voice.NewAgent(msgBus, transcriber)
|
||||
voiceAgent.Start(vaCtx)
|
||||
}
|
||||
|
||||
enabledChannels := runningServices.ChannelManager.GetEnabledChannels()
|
||||
|
|
@ -335,6 +343,9 @@ func stopAndCleanupServices(runningServices *services, shutdownTimeout time.Dura
|
|||
if !isReload && runningServices.ChannelManager != nil {
|
||||
runningServices.ChannelManager.StopAll(shutdownCtx)
|
||||
}
|
||||
if runningServices.VoiceAgentCancel != nil {
|
||||
runningServices.VoiceAgentCancel()
|
||||
}
|
||||
if runningServices.DeviceService != nil {
|
||||
runningServices.DeviceService.Stop()
|
||||
}
|
||||
|
|
@ -515,10 +526,16 @@ func restartServices(
|
|||
fmt.Println(" ✓ Device event service restarted")
|
||||
}
|
||||
|
||||
transcriber := voice.DetectTranscriber(cfg)
|
||||
transcriber := asr.DetectTranscriber(cfg)
|
||||
al.SetTranscriber(transcriber)
|
||||
if transcriber != nil {
|
||||
logger.InfoCF("voice", "Transcription re-enabled (agent-level)", map[string]any{"provider": transcriber.Name()})
|
||||
|
||||
// Start Voice Agent Orchestrator on reload
|
||||
vaCtx, vaCancel := context.WithCancel(context.Background())
|
||||
runningServices.VoiceAgentCancel = vaCancel
|
||||
voiceAgent := voice.NewAgent(msgBus, transcriber)
|
||||
voiceAgent.Start(vaCtx)
|
||||
} else {
|
||||
logger.InfoCF("voice", "Transcription disabled", nil)
|
||||
}
|
||||
|
|
|
|||
144
pkg/tts/tts.go
Normal file
144
pkg/tts/tts.go
Normal file
|
|
@ -0,0 +1,144 @@
|
|||
package tts
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
)
|
||||
|
||||
type TTSProvider interface {
|
||||
Name() string
|
||||
Synthesize(ctx context.Context, text string) (io.ReadCloser, error)
|
||||
}
|
||||
|
||||
type OpenAITTSProvider struct {
|
||||
apiKey string
|
||||
apiBase string
|
||||
voice string
|
||||
model string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func NewOpenAITTSProvider(apiKey string, apiBase string, proxyURL string) *OpenAITTSProvider {
|
||||
// Normalize apiBase to avoid malformed endpoints like
|
||||
// "https://api.openai.com/audio/speech" when "/v1" is required.
|
||||
if apiBase == "" {
|
||||
apiBase = "https://api.openai.com/v1/audio/speech"
|
||||
} else {
|
||||
if u, err := url.Parse(apiBase); err == nil && u.Scheme != "" && u.Host != "" {
|
||||
path := u.Path
|
||||
if u.Host == "api.openai.com" {
|
||||
// For the official OpenAI host, ensure exactly one /v1 prefix and
|
||||
// that the path ends with /audio/speech.
|
||||
if path == "" || path == "/" || path == "/v1" {
|
||||
path = "/v1/audio/speech"
|
||||
} else {
|
||||
if !strings.HasPrefix(path, "/") {
|
||||
path = "/" + path
|
||||
}
|
||||
if !strings.HasPrefix(path, "/v1/") {
|
||||
path = "/v1" + strings.TrimSuffix(path, "/")
|
||||
}
|
||||
if !strings.HasSuffix(path, "/audio/speech") {
|
||||
path = strings.TrimSuffix(path, "/") + "/audio/speech"
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// For non-OpenAI hosts (e.g., proxies), preserve the existing base
|
||||
// path and only ensure it ends with /audio/speech.
|
||||
if !strings.HasSuffix(path, "/audio/speech") {
|
||||
path = strings.TrimSuffix(path, "/") + "/audio/speech"
|
||||
}
|
||||
}
|
||||
u.Path = path
|
||||
apiBase = u.String()
|
||||
} else {
|
||||
// Fallback to the previous string-based behavior if parsing fails.
|
||||
if apiBase == "https://api.openai.com/v1" {
|
||||
apiBase = "https://api.openai.com/v1/audio/speech"
|
||||
} else if !strings.HasSuffix(apiBase, "/audio/speech") {
|
||||
// Just in case they provide openrouter base or standard base
|
||||
apiBase = strings.TrimSuffix(apiBase, "/") + "/audio/speech"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
client := &http.Client{
|
||||
Timeout: 60 * time.Second,
|
||||
}
|
||||
|
||||
if proxyURL != "" {
|
||||
if pURL, err := url.Parse(proxyURL); err == nil {
|
||||
client.Transport = &http.Transport{
|
||||
Proxy: http.ProxyURL(pURL),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return &OpenAITTSProvider{
|
||||
apiKey: apiKey,
|
||||
apiBase: apiBase,
|
||||
voice: "alloy",
|
||||
model: "tts-1",
|
||||
httpClient: client,
|
||||
}
|
||||
}
|
||||
|
||||
func (t *OpenAITTSProvider) Name() string {
|
||||
return "openai-tts"
|
||||
}
|
||||
|
||||
func (t *OpenAITTSProvider) Synthesize(ctx context.Context, text string) (io.ReadCloser, error) {
|
||||
logger.InfoCF("voice-tts", "Starting TTS synthesis", map[string]any{"text_len": len(text)})
|
||||
|
||||
reqBody := map[string]any{
|
||||
"model": t.model,
|
||||
"input": text,
|
||||
"voice": t.voice,
|
||||
"response_format": "opus",
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", t.apiBase, bytes.NewReader(jsonData))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+t.apiKey)
|
||||
|
||||
resp, err := t.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to send request: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
defer resp.Body.Close()
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("API error (status %d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
return resp.Body, nil
|
||||
}
|
||||
|
||||
func DetectTTS(cfg *config.Config) TTSProvider {
|
||||
for _, mc := range cfg.ModelList {
|
||||
if strings.Contains(strings.ToLower(mc.ModelName), "tts") && mc.APIKey != "" {
|
||||
return NewOpenAITTSProvider(mc.APIKey, mc.APIBase, mc.Proxy)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
253
pkg/voice/agent.go
Normal file
253
pkg/voice/agent.go
Normal file
|
|
@ -0,0 +1,253 @@
|
|||
package voice
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/pion/rtp"
|
||||
"github.com/pion/webrtc/v3/pkg/media/oggwriter"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/asr"
|
||||
"github.com/sipeed/picoclaw/pkg/bus"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
)
|
||||
|
||||
type speechAccumulator struct {
|
||||
writer *oggwriter.OggWriter
|
||||
file string
|
||||
lastAudioAt time.Time
|
||||
mu sync.Mutex
|
||||
closed bool
|
||||
chatID string
|
||||
speakerID string
|
||||
sessionID string
|
||||
channel string
|
||||
}
|
||||
|
||||
func (a *speechAccumulator) Push(chunk bus.AudioChunk) {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
if a.closed {
|
||||
return
|
||||
}
|
||||
|
||||
a.lastAudioAt = time.Now()
|
||||
|
||||
pkt := &rtp.Packet{
|
||||
Header: rtp.Header{
|
||||
SequenceNumber: uint16(chunk.Sequence),
|
||||
Timestamp: chunk.Timestamp,
|
||||
SSRC: uint32(chunk.Sequence), // Arbitrary dummy
|
||||
},
|
||||
Payload: chunk.Data,
|
||||
}
|
||||
|
||||
if err := a.writer.WriteRTP(pkt); err != nil {
|
||||
logger.ErrorCF("voice-agent", "Failed to write RTP", map[string]any{"error": err})
|
||||
}
|
||||
}
|
||||
|
||||
func (a *speechAccumulator) Close() {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
if !a.closed {
|
||||
a.writer.Close()
|
||||
a.closed = true
|
||||
}
|
||||
}
|
||||
|
||||
type Agent struct {
|
||||
bus *bus.MessageBus
|
||||
transcriber asr.Transcriber
|
||||
|
||||
mu sync.Mutex
|
||||
sessions map[string]*speechAccumulator // keyed by sessionID_speakerID
|
||||
}
|
||||
|
||||
func NewAgent(mb *bus.MessageBus, t asr.Transcriber) *Agent {
|
||||
return &Agent{
|
||||
bus: mb,
|
||||
transcriber: t,
|
||||
sessions: make(map[string]*speechAccumulator),
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) Start(ctx context.Context) {
|
||||
logger.InfoCF("voice-agent", "Started Voice Agent orchestrator", nil)
|
||||
go a.listenChunks(ctx)
|
||||
go a.vadTick(ctx)
|
||||
|
||||
// Cleanup sessions on shutdown
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
a.mu.Lock()
|
||||
for key, acc := range a.sessions {
|
||||
acc.Close()
|
||||
os.Remove(acc.file)
|
||||
delete(a.sessions, key)
|
||||
}
|
||||
a.mu.Unlock()
|
||||
logger.InfoCF("voice-agent", "Cleaned up voice sessions on shutdown", nil)
|
||||
}()
|
||||
}
|
||||
|
||||
func (a *Agent) listenChunks(ctx context.Context) {
|
||||
chunks := a.bus.AudioChunksChan()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case chunk, ok := <-chunks:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
a.handleChunk(chunk)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) handleChunk(chunk bus.AudioChunk) {
|
||||
// Only accept Opus-encoded audio
|
||||
if chunk.Format != "opus" {
|
||||
logger.DebugCF("voice-agent", "Ignoring unsupported audio format", map[string]any{"format": chunk.Format})
|
||||
return
|
||||
}
|
||||
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
key := fmt.Sprintf("%s_%s", chunk.SessionID, chunk.SpeakerID)
|
||||
|
||||
acc, exists := a.sessions[key]
|
||||
if !exists {
|
||||
filename := filepath.Join(os.TempDir(), fmt.Sprintf("voice_%s_%d.ogg", key, time.Now().UnixNano()))
|
||||
writer, err := oggwriter.New(filename, uint32(chunk.SampleRate), uint16(chunk.Channels))
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice-agent", "Failed to create OggWriter", map[string]any{"error": err})
|
||||
return
|
||||
}
|
||||
|
||||
acc = &speechAccumulator{
|
||||
writer: writer,
|
||||
file: filename,
|
||||
lastAudioAt: time.Now(),
|
||||
chatID: chunk.ChatID,
|
||||
speakerID: chunk.SpeakerID,
|
||||
sessionID: chunk.SessionID,
|
||||
channel: chunk.Channel,
|
||||
}
|
||||
a.sessions[key] = acc
|
||||
logger.DebugCF("voice-agent", "Started accumulating voice", map[string]any{"key": key, "file": filename})
|
||||
}
|
||||
|
||||
acc.Push(chunk)
|
||||
}
|
||||
|
||||
func (a *Agent) vadTick(ctx context.Context) {
|
||||
ticker := time.NewTicker(500 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
a.checkSilence(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) checkSilence(ctx context.Context) {
|
||||
a.mu.Lock()
|
||||
now := time.Now()
|
||||
var finished []*speechAccumulator
|
||||
|
||||
for key, acc := range a.sessions {
|
||||
acc.mu.Lock()
|
||||
last := acc.lastAudioAt
|
||||
acc.mu.Unlock()
|
||||
|
||||
if now.Sub(last) > 1500*time.Millisecond {
|
||||
acc.Close()
|
||||
delete(a.sessions, key)
|
||||
finished = append(finished, acc)
|
||||
}
|
||||
}
|
||||
a.mu.Unlock()
|
||||
|
||||
for _, acc := range finished {
|
||||
go a.processUtterance(ctx, acc)
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) processUtterance(ctx context.Context, acc *speechAccumulator) {
|
||||
defer os.Remove(acc.file)
|
||||
|
||||
logger.InfoCF("voice-agent", "User finished speaking, transcribing...", map[string]any{"file": acc.file})
|
||||
|
||||
if a.transcriber == nil {
|
||||
logger.ErrorCF("voice-agent", "No STT configured!", nil)
|
||||
return
|
||||
}
|
||||
|
||||
res, err := a.transcriber.Transcribe(ctx, acc.file)
|
||||
if err != nil {
|
||||
logger.ErrorCF("voice-agent", "Transcription failed", map[string]any{"error": err})
|
||||
return
|
||||
}
|
||||
|
||||
if res.Text == "" {
|
||||
logger.DebugCF("voice-agent", "Ignored empty transcription", map[string]any{"file": acc.file})
|
||||
return
|
||||
}
|
||||
|
||||
logger.InfoCF("voice-agent", "Transcription result", map[string]any{"text": res.Text, "duration": res.Duration})
|
||||
|
||||
channelType := acc.channel
|
||||
if channelType == "" {
|
||||
channelType = "discord" // fallback for legacy chunks
|
||||
}
|
||||
|
||||
text := strings.ToLower(strings.TrimSpace(res.Text))
|
||||
if strings.Contains(text, "leave the voice channel") || strings.Contains(text, "leave voice") ||
|
||||
strings.Contains(text, "disconnect voice") || strings.Contains(text, "leave the channel") ||
|
||||
strings.Contains(text, "leave channel") {
|
||||
logger.InfoCF("voice-agent", "Voice command triggered: leave", nil)
|
||||
if err := a.bus.PublishVoiceControl(ctx, bus.VoiceControl{
|
||||
SessionID: acc.sessionID,
|
||||
Type: "command",
|
||||
Action: "leave",
|
||||
}); err != nil {
|
||||
logger.ErrorCF("voice-agent", "Failed to publish leave control", map[string]any{"error": err})
|
||||
}
|
||||
if err := a.bus.PublishOutbound(ctx, bus.OutboundMessage{
|
||||
Channel: channelType,
|
||||
ChatID: acc.chatID,
|
||||
Content: "Goodbye! Leaving the voice channel.",
|
||||
}); err != nil {
|
||||
logger.ErrorCF("voice-agent", "Failed to publish goodbye message", map[string]any{"error": err})
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
oralPrompt := "\n\n[SYSTEM]: The user just spoke this to you over voice chat. Please reply in a highly concise, conversational, oral style suitable for text-to-speech. Do not use markdown, emojis, asterisks, or code blocks. Speak naturally."
|
||||
|
||||
if err := a.bus.PublishInbound(ctx, bus.InboundMessage{
|
||||
Channel: channelType,
|
||||
SenderID: acc.speakerID,
|
||||
ChatID: acc.chatID,
|
||||
Content: res.Text + oralPrompt,
|
||||
Peer: bus.Peer{Kind: "channel", ID: acc.chatID},
|
||||
Metadata: map[string]string{
|
||||
"is_voice": "true",
|
||||
},
|
||||
}); err != nil {
|
||||
logger.ErrorCF("voice-agent", "Failed to publish inbound message", map[string]any{"error": err})
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue