feat: merge pipeline, per-user credentials, unified picker, group contacts

- Enable merge UI for linking channel contacts to tenant_users
- Contact → tenant_user resolution with cached lookup (60s TTL)
- MCP per-user credentials via user-keyed connection pool
- Secure CLI per-user credentials with AES-256-GCM encryption
- Unified UserPickerCombobox searching contacts + tenant_users
- Group contact collection with chat title in all channels
- Group permission inheritance via wildcard user_id="*"
- Fix heartbeat using wrong userID in group chats
- Filter internal senders from contact collection
- Add contact_type column (user/group) to channel_contacts
- SQLite schema v2 migration for desktop edition
This commit is contained in:
viettranx committed 2026-03-29 22:33:17 +07:00
1 parent 4cb6991fc6
commit 21b6c454ca
81 files changed
+2015 -632

No files matched your search

+29 -2
View File
@@ -109,7 +109,8 @@ func processNormalMessage(
}
// Auto-collect channel contacts for the contact selector.
if deps.ContactCollector != nil && msg.SenderID != "" {
// Skip internal senders (system:*, notification:*, teammate:*, ticker:*, session_send_tool).
if deps.ContactCollector != nil && msg.SenderID != "" && !bus.IsInternalSender(msg.SenderID) {
senderNumericID := msg.SenderID
if idx := strings.IndexByte(senderNumericID, '|'); idx > 0 {
senderNumericID = senderNumericID[:idx]
@@ -120,7 +121,33 @@ func processNormalMessage(
}
displayName := sessionMeta["display_name"]
username := sessionMeta["username"]
deps.ContactCollector.EnsureContact(ctx, channelType, msg.Channel, senderNumericID, userID, displayName, username, peerKind)
deps.ContactCollector.EnsureContact(ctx, channelType, msg.Channel, senderNumericID, userID, displayName, username, peerKind, "user")
// Also collect group chat as a contact (for group permission management / merge).
// Group IDs (e.g., Telegram "-100456") differ from user IDs — no UNIQUE conflict.
if peerKind == string(sessions.PeerGroup) && msg.ChatID != "" {
groupTitle := msg.Metadata["chat_title"] // Telegram: message.Chat.Title
deps.ContactCollector.EnsureContact(ctx, channelType, msg.Channel, msg.ChatID, "", groupTitle, "", "group", "group")
}
}
// --- Resolve merged tenant user identity ---
// If the sender has been merged to a tenant_user, use the tenant user's ID
// for DM sessions. This enables per-user features (MCP creds, SecureCLI creds).
// Group sessions keep the group-scoped userID; sender resolution happens via SenderID.
if deps.ContactCollector != nil && peerKind == string(sessions.PeerDirect) && msg.SenderID != "" && !bus.IsInternalSender(msg.SenderID) {
senderNumeric := msg.SenderID
if idx := strings.IndexByte(senderNumeric, '|'); idx > 0 {
senderNumeric = senderNumeric[:idx]
}
chType := deps.ChannelMgr.ChannelTypeForName(msg.Channel)
if chType == "" {
chType = msg.Channel
}
if resolved, err := deps.ContactCollector.ResolveTenantUserID(ctx, chType, senderNumeric); err == nil && resolved != "" {
slog.Debug("contact.resolved_tenant_user", "sender", senderNumeric, "tenant_user", resolved)
userID = resolved
}
}
// --- Quota check ---
+6 -4
View File
@@ -203,6 +203,11 @@ func (l *Loop) runLoop(ctx context.Context, req RunRequest) (*RunResult, error)
// channel type, and final-iteration stripping.
var toolDefs []providers.ToolDefinition
var allowedTools map[string]bool
// Resolve per-user MCP tools (servers requiring user credentials).
// Must run before buildFilteredTools so tools are in the Registry for policy filtering.
if req.UserID != "" {
l.getUserMCPTools(iterCtx, req.UserID)
}
toolDefs, allowedTools, messages = l.buildFilteredTools(&req, hadBootstrap, rs.iteration, maxIter, messages)
// Use per-request overrides if set (e.g. heartbeat uses cheaper provider/model).
@@ -321,10 +326,7 @@ func (l *Loop) runLoop(ctx context.Context, req RunRequest) (*RunResult, error)
// Calibrate overhead on first LLM response with usage data.
if !rs.overheadCalibrated && resp.Usage != nil && resp.Usage.PromptTokens > 0 {
historyEst := EstimateHistoryTokens(messages)
rs.overheadTokens = resp.Usage.PromptTokens - historyEst
if rs.overheadTokens < 0 {
rs.overheadTokens = 0
}
rs.overheadTokens = max(resp.Usage.PromptTokens-historyEst, 0)
rs.overheadCalibrated = true
}
+1 -1
View File
@@ -22,7 +22,7 @@ func isUserFilePopulated(content string) bool {
return false
}
// Template markers: "**Name:**" followed by newline (no value) or just whitespace
for _, line := range strings.Split(content, "\n") {
for line := range strings.SplitSeq(content, "\n") {
line = strings.TrimSpace(line)
if line == "- **Name:**" || line == "**Name:**" {
return false // name field still empty
+4 -16
View File
@@ -547,10 +547,7 @@ func (l *Loop) maybeSummarize(ctx context.Context, sessionKey string) {
// lastPromptTokens includes everything (system prompt, tools, context files, history).
// We subtract estimated overhead so the threshold comparison is history-only.
lastPT, lastMC := l.sessions.GetLastPromptTokens(ctx, sessionKey)
adjustedLastPT := lastPT - l.estimateOverhead(history, lastPT, lastMC)
if adjustedLastPT < 0 {
adjustedLastPT = 0
}
adjustedLastPT := max(lastPT-l.estimateOverhead(history, lastPT, lastMC), 0)
tokenEstimate := EstimateTokensWithCalibration(history, adjustedLastPT, lastMC)
// Resolve compaction threshold from config: token-only (no message count guard).
@@ -666,23 +663,14 @@ func (l *Loop) maybeSummarize(ctx context.Context, sessionKey string) {
func (l *Loop) estimateOverhead(history []providers.Message, lastPromptTokens, lastMsgCount int) int {
if lastPromptTokens <= 0 || lastMsgCount <= 0 {
// No calibration data — use conservative default (20% of context, capped at 40k).
fallback := int(float64(l.contextWindow) * 0.2)
if fallback > 40000 {
fallback = 40000
}
fallback := min(int(float64(l.contextWindow)*0.2), 40000)
return fallback
}
// Overhead = total prompt tokens - estimated history tokens at calibration time.
count := lastMsgCount
if count > len(history) {
count = len(history)
}
count := min(lastMsgCount, len(history))
historyEstAtCalibration := EstimateHistoryTokens(history[:count])
overhead := lastPromptTokens - historyEstAtCalibration
if overhead < 0 {
overhead = 0
}
overhead := max(lastPromptTokens-historyEstAtCalibration, 0)
// Clamp: overhead shouldn't exceed 40% of context window.
maxOverhead := int(float64(l.contextWindow) * 0.4)
if overhead > maxOverhead {
+1 -1
View File
@@ -217,7 +217,7 @@ func TestSanitizeHistory_DedupAcrossTwoTurns(t *testing.T) {
func TestSanitizeHistory_LargeHistory_Performance(t *testing.T) {
// Build 1000-message history with proper tool pairing
msgs := make([]providers.Message, 0, 1000)
for i := 0; i < 250; i++ {
for i := range 250 {
tcID := "tc_" + strings.Repeat("x", 5) + "_" + string(rune('a'+i%26)) + string(rune('0'+i%10))
msgs = append(msgs,
providers.Message{Role: "user", Content: "question " + tcID},
+106
View File
@@ -0,0 +1,106 @@
package agent
import (
"context"
"log/slog"
"maps"
mcpbridge "github.com/nextlevelbuilder/goclaw/internal/mcp"
"github.com/nextlevelbuilder/goclaw/internal/tools"
)
// getUserMCPTools returns per-user MCP tools for servers requiring user credentials.
// Tools are cached per-user in mcpUserTools sync.Map and registered in the shared
// tool registry so ExecuteWithContext can resolve them. On first call for a user,
// connections are established via pool.AcquireUser() and BridgeTools created.
func (l *Loop) getUserMCPTools(ctx context.Context, userID string) []tools.Tool {
if len(l.mcpUserCredSrvs) == 0 || l.mcpPool == nil || l.mcpStore == nil || userID == "" {
return nil
}
if cached, ok := l.mcpUserTools.Load(userID); ok {
cachedTools := cached.([]tools.Tool)
// Check if any cached tool's connection was evicted by pool.
// If so, clear cache and re-acquire connections.
allConnected := true
for _, t := range cachedTools {
if bt, ok := t.(interface{ IsConnected() bool }); ok && !bt.IsConnected() {
allConnected = false
break
}
}
if allConnected {
return cachedTools
}
l.mcpUserTools.Delete(userID)
slog.Debug("mcp.user_tools_stale", "user", userID, "reason", "pool_evicted")
}
var userTools []tools.Tool
for _, info := range l.mcpUserCredSrvs {
srv := info.Server
// Check if user has credentials for this server
uc, err := l.mcpStore.GetUserCredentials(ctx, srv.ID, userID)
if err != nil || uc == nil || (uc.APIKey == "" && len(uc.Headers) == 0 && len(uc.Env) == 0) {
continue
}
// Resolve connection params: server defaults merged with user overrides
args := mcpbridge.ParseJSONBytesToStringSlice(srv.Args)
env := mcpbridge.ParseJSONBytesToStringMap(srv.Env)
if env == nil {
env = make(map[string]string)
}
headers := mcpbridge.ParseJSONBytesToStringMap(srv.Headers)
if headers == nil {
headers = make(map[string]string)
}
// Inject server-level API key into headers if present
if srv.APIKey != "" && headers["Authorization"] == "" {
headers["Authorization"] = "Bearer " + srv.APIKey
}
// Merge user credentials (user overrides server defaults)
if uc.APIKey != "" {
headers["Authorization"] = "Bearer " + uc.APIKey
}
maps.Copy(headers, uc.Headers)
maps.Copy(env, uc.Env)
// Acquire user-keyed pool connection
entry, err := l.mcpPool.AcquireUser(ctx, l.tenantID, srv.Name, userID,
srv.Transport, srv.Command, args, env, srv.URL, headers, srv.TimeoutSec)
if err != nil {
slog.Warn("mcp.user_pool_acquire_failed", "server", srv.Name, "user", userID, "error", err)
continue
}
// Release immediately — BridgeTools hold client pointer directly.
// This allows pool idle eviction to work (refCount=0 + lastUsed for TTL).
// When pool evicts the connection, BridgeTool.Execute detects connected=false.
l.mcpPool.ReleaseUser(mcpbridge.UserPoolKey(l.tenantID, srv.Name, userID))
// Create BridgeTools pointing to user's connection and register in the
// shared tool registry so ExecuteWithContext can resolve them by name.
reg, _ := l.tools.(*tools.Registry)
for _, mcpTool := range entry.MCPTools() {
bt := mcpbridge.NewBridgeTool(srv.Name, mcpTool, entry.Client(), srv.ToolPrefix, srv.TimeoutSec, entry.Connected())
// Register in registry so ExecuteWithContext can find them.
// Skip if already registered (another user loaded this server with same tool names).
if reg != nil {
if _, exists := reg.Get(bt.Name()); !exists {
reg.Register(bt)
}
}
userTools = append(userTools, bt)
}
}
if len(userTools) > 0 {
l.mcpUserTools.Store(userID, userTools)
slog.Info("mcp.user_tools_loaded", "user", userID, "tools", len(userTools))
}
return userTools
}
+2
View File
@@ -10,6 +10,8 @@ import (
// buildFilteredTools resolves the per-iteration tool definitions based on policy,
// disabled tools, bootstrap mode, skill visibility, channel type, and iteration budget.
// Per-user MCP tools must be registered in the Registry before calling this function
// (via getUserMCPTools) so they are included in policy filtering and execution.
// Returns tool definitions for the provider, an allowed-tools map for execution validation,
// and the (potentially modified) messages slice when final-iteration stripping appends a hint.
func (l *Loop) buildFilteredTools(req *RunRequest, hadBootstrap bool, iteration, maxIter int, messages []providers.Message) ([]providers.ToolDefinition, map[string]bool, []providers.Message) {
+15
View File
@@ -10,6 +10,7 @@ import (
"github.com/nextlevelbuilder/goclaw/internal/bootstrap"
"github.com/nextlevelbuilder/goclaw/internal/bus"
"github.com/nextlevelbuilder/goclaw/internal/config"
mcpbridge "github.com/nextlevelbuilder/goclaw/internal/mcp"
"github.com/nextlevelbuilder/goclaw/internal/media"
"github.com/nextlevelbuilder/goclaw/internal/providers"
"github.com/nextlevelbuilder/goclaw/internal/sandbox"
@@ -110,6 +111,12 @@ type Loop struct {
cacheInvalidate CacheInvalidateFunc // invalidate context file cache after seeding
userSetups sync.Map // userID → *userSetup (workspace + seeding state, per Loop instance)
// Per-user MCP tools: servers requiring user credentials get connected per-request.
mcpStore store.MCPServerStore // for credential lookup
mcpPool *mcpbridge.Pool // user-keyed connection pool
mcpUserCredSrvs []store.MCPAccessInfo // servers needing per-user creds
mcpUserTools sync.Map // userID → []tools.Tool (cached per-user tools)
// Compaction config (memory flush settings)
compactionCfg *config.CompactionConfig
@@ -304,6 +311,11 @@ type LoopConfig struct {
// Memory store for extractive memory fallback (writes directly when LLM flush fails)
MemoryStore store.MemoryStore
// Per-user MCP tools (servers requiring per-user credentials)
MCPStore store.MCPServerStore // for credential lookup
MCPPool *mcpbridge.Pool // user-keyed connection pool
MCPUserCredSrvs []store.MCPAccessInfo // servers needing per-user creds
}
const defaultMaxTokens = config.DefaultMaxTokens
@@ -398,6 +410,9 @@ func NewLoop(cfg LoopConfig) *Loop {
budgetMonthlyCents: cfg.BudgetMonthlyCents,
tracingStore: cfg.TracingStore,
memStore: cfg.MemoryStore,
mcpStore: cfg.MCPStore,
mcpPool: cfg.MCPPool,
mcpUserCredSrvs: cfg.MCPUserCredSrvs,
}
}
+19 -12
View File
@@ -256,6 +256,7 @@ func NewManagedResolver(deps ResolverDeps) ResolverFunc {
// one agent pollute the shared deps. Tools and become visible to ALL agents
// (even those without MCP grants), because FilterTools reads from registry.List().
hasMCPTools := false
var mcpUserCredSrvs []store.MCPAccessInfo
if deps.MCPStore != nil {
if toolsReg == deps.Tools {
toolsReg = deps.Tools.Clone()
@@ -268,20 +269,23 @@ func NewManagedResolver(deps ResolverDeps) ResolverFunc {
mcpMgr := mcpbridge.NewManager(toolsReg, mcpOpts...)
if err := mcpMgr.LoadForAgent(ctx, ag.ID, ""); err != nil {
slog.Warn("failed to load MCP servers for agent", "agent", agentKey, "error", err)
} else if mcpMgr.IsSearchMode() {
// Search mode: too many tools — register mcp_tool_search meta-tool.
// Also wire lazy activator so deferred tools can be called by name directly.
toolsReg.SetDeferredActivator(mcpMgr.ActivateToolIfDeferred)
searchTool := mcpbridge.NewMCPToolSearchTool(mcpMgr)
toolsReg.Register(searchTool)
hasMCPTools = true
slog.Info("mcp.agent.search_mode", "agent", agentKey,
"deferred_tools", len(mcpMgr.DeferredToolInfos()))
} else {
toolNames := mcpMgr.ToolNames()
if len(toolNames) > 0 {
mcpUserCredSrvs = mcpMgr.UserCredServers()
if mcpMgr.IsSearchMode() {
// Search mode: too many tools — register mcp_tool_search meta-tool.
// Also wire lazy activator so deferred tools can be called by name directly.
toolsReg.SetDeferredActivator(mcpMgr.ActivateToolIfDeferred)
searchTool := mcpbridge.NewMCPToolSearchTool(mcpMgr)
toolsReg.Register(searchTool)
hasMCPTools = true
slog.Info("mcp.agent.tools_loaded", "agent", agentKey, "tools", len(toolNames))
slog.Info("mcp.agent.search_mode", "agent", agentKey,
"deferred_tools", len(mcpMgr.DeferredToolInfos()))
} else {
toolNames := mcpMgr.ToolNames()
if len(toolNames) > 0 {
hasMCPTools = true
slog.Info("mcp.agent.tools_loaded", "agent", agentKey, "tools", len(toolNames))
}
}
}
}
@@ -398,6 +402,9 @@ func NewManagedResolver(deps ResolverDeps) ResolverFunc {
BudgetMonthlyCents: derefInt(ag.BudgetMonthlyCents),
TracingStore: deps.TracingStore,
MemoryStore: deps.MemoryStore,
MCPStore: deps.MCPStore,
MCPPool: deps.MCPPool,
MCPUserCredSrvs: mcpUserCredSrvs,
})
slog.Info("resolved agent from DB", "agent", agentKey, "model", ag.Model, "provider", ag.Provider)
+1 -1
View File
@@ -15,7 +15,7 @@ import (
// retryOnBusy retries fn up to 3 times on SQLITE_BUSY errors with 500ms delay.
func retryOnBusy(fn func() error) error {
var lastErr error
for attempt := 0; attempt < 3; attempt++ {
for attempt := range 3 {
lastErr = fn()
if lastErr == nil {
return nil
+4 -6
View File
@@ -221,9 +221,7 @@ func TestBroadcast_ConcurrentSubscribeUnsubscribe(t *testing.T) {
done := make(chan struct{})
// Broadcast in a goroutine
wg.Add(1)
go func() {
defer wg.Done()
wg.Go(func() {
for {
select {
case <-done:
@@ -232,10 +230,10 @@ func TestBroadcast_ConcurrentSubscribeUnsubscribe(t *testing.T) {
mb.Broadcast(Event{Name: "concurrent"})
}
}
}()
})
// Subscribe/unsubscribe rapidly
for i := 0; i < 100; i++ {
for range 100 {
mb.Subscribe("rapid", func(e Event) {})
mb.Unsubscribe("rapid")
}
@@ -252,7 +250,7 @@ func TestPublishInbound_ConcurrentProducers(t *testing.T) {
const n = 100
var wg sync.WaitGroup
wg.Add(n)
for i := 0; i < n; i++ {
for range n {
go func() {
defer wg.Done()
mb.TryPublishInbound(InboundMessage{Content: "msg"})
+11
View File
@@ -3,6 +3,7 @@ package bus
import (
"context"
"encoding/json"
"strings"
"github.com/google/uuid"
)
@@ -160,3 +161,13 @@ type MessageRouter interface {
PublishOutbound(msg OutboundMessage)
SubscribeOutbound(ctx context.Context) (OutboundMessage, bool)
}
// IsInternalSender returns true if the senderID belongs to an internal system
// component (not a real channel user). These should not be stored as contacts.
func IsInternalSender(senderID string) bool {
return strings.HasPrefix(senderID, "system:") ||
strings.HasPrefix(senderID, "notification:") ||
strings.HasPrefix(senderID, "teammate:") ||
strings.HasPrefix(senderID, "ticker:") ||
senderID == "session_send_tool"
}
+2 -2
View File
@@ -195,7 +195,7 @@ func (c *Channel) handleMessage(_ *discordgo.Session, m *discordgo.MessageCreate
// Collect contact even when bot is not mentioned (cache prevents DB spam).
if cc := c.ContactCollector(); cc != nil {
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, senderName, m.Author.Username, "group")
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, senderName, m.Author.Username, "group", "user")
}
slog.Debug("discord group message recorded (no mention)",
@@ -287,7 +287,7 @@ func (c *Channel) handleMessage(_ *discordgo.Session, m *discordgo.MessageCreate
// Collect contact for processed messages (DM + group-mentioned).
if cc := c.ContactCollector(); cc != nil {
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, senderName, m.Author.Username, peerKind)
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, senderName, m.Author.Username, peerKind, "user")
}
// Publish directly to bus (to preserve MediaFile MIME types)
+2 -2
View File
@@ -106,7 +106,7 @@ func (c *Channel) handleMessageEvent(ctx context.Context, event *MessageEvent) {
// Collect contact even when bot is not mentioned (cache prevents DB spam).
if cc := c.ContactCollector(); cc != nil {
cc.EnsureContact(ctx, c.Type(), c.Name(), mc.SenderID, mc.SenderID, senderName, "", "group")
cc.EnsureContact(ctx, c.Type(), c.Name(), mc.SenderID, mc.SenderID, senderName, "", "group", "user")
}
slog.Debug("feishu group message recorded (no mention)",
@@ -162,7 +162,7 @@ func (c *Channel) handleMessageEvent(ctx context.Context, event *MessageEvent) {
// Collect contact for processed messages (DM + group-mentioned).
if cc := c.ContactCollector(); cc != nil {
cc.EnsureContact(ctx, c.Type(), c.Name(), mc.SenderID, mc.SenderID, senderName, "", peerKind)
cc.EnsureContact(ctx, c.Type(), c.Name(), mc.SenderID, mc.SenderID, senderName, "", peerKind, "user")
}
metadata := map[string]string{
+1 -1
View File
@@ -575,7 +575,7 @@ func (c *Channel) ListGroupMembers(ctx context.Context, chatID string) ([]channe
}
// Auto-sync member into contact store
if cc := c.ContactCollector(); cc != nil {
cc.EnsureContact(ctx, channels.TypeFeishu, c.Name(), m.MemberID, m.MemberID, m.Name, "", "group")
cc.EnsureContact(ctx, channels.TypeFeishu, c.Name(), m.MemberID, m.MemberID, m.Name, "", "group", "user")
}
}
return result, nil
+1 -1
View File
@@ -200,7 +200,7 @@ func (c *Channel) handleMessage(ev *slackevents.MessageEvent) {
// Collect contact even when bot is not mentioned (cache prevents DB spam).
if cc := c.ContactCollector(); cc != nil {
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, displayName, "", "group")
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, displayName, "", "group", "user")
}
slog.Debug("slack group message recorded (no mention)",
+1 -1
View File
@@ -34,7 +34,7 @@ func (c *Channel) HandleMessage(senderID, chatID, content string, mediaPaths []s
// Collect contact for processed messages (DM + group-mentioned).
if cc := c.ContactCollector(); cc != nil {
ctx := store.WithTenantID(context.Background(), c.TenantID())
cc.EnsureContact(ctx, c.Type(), c.Name(), userID, userID, metadata["username"], "", peerKind)
cc.EnsureContact(ctx, c.Type(), c.Name(), userID, userID, metadata["username"], "", peerKind, "user")
}
c.Bus().PublishInbound(bus.InboundMessage{
+8 -2
View File
@@ -331,7 +331,9 @@ func (c *Channel) handleMessage(ctx context.Context, update telego.Update) {
// Collect contact even when bot is not mentioned (cache prevents DB spam).
if cc := c.ContactCollector(); cc != nil {
contactName := strings.TrimSpace(user.FirstName + " " + user.LastName)
cc.EnsureContact(ctx, c.Type(), c.Name(), userID, userID, contactName, user.Username, "group")
cc.EnsureContact(ctx, c.Type(), c.Name(), userID, userID, contactName, user.Username, "group", "user")
// Also collect group chat itself as a contact (for group permission / merge).
cc.EnsureContact(ctx, c.Type(), c.Name(), chatIDStr, "", message.Chat.Title, "", "group", "group")
}
slog.Debug("telegram group message recorded (no mention)",
@@ -583,7 +585,11 @@ func (c *Channel) handleMessage(ctx context.Context, update telego.Update) {
// Collect contact for processed messages (DM + group-mentioned).
if cc := c.ContactCollector(); cc != nil {
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, userID, user.FirstName, user.Username, peerKind)
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, userID, user.FirstName, user.Username, peerKind, "user")
// Also collect group chat itself as a contact (for group permission / merge).
if isGroup {
cc.EnsureContact(ctx, c.Type(), c.Name(), chatIDStr, "", message.Chat.Title, "", "group", "group")
}
}
c.Bus().PublishInbound(bus.InboundMessage{
+1 -1
View File
@@ -271,7 +271,7 @@ func (c *Channel) handleIncomingMessage(msg map[string]any) {
// Collect contact for processed messages.
if cc := c.ContactCollector(); cc != nil {
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, metadata["user_name"], "", peerKind)
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, metadata["user_name"], "", peerKind, "user")
}
c.HandleMessage(senderID, chatID, content, media, metadata, peerKind)
+3 -3
View File
@@ -64,7 +64,7 @@ func (c *Channel) handleDM(msg protocol.UserMessage) {
// Collect contact for DM messages.
if cc := c.ContactCollector(); cc != nil {
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, senderName, "", "direct")
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, senderName, "", "direct", "user")
}
metadata := map[string]string{
@@ -111,7 +111,7 @@ func (c *Channel) handleGroupMessage(msg protocol.GroupMessage) {
// Collect contact even when bot is not mentioned (cache prevents DB spam).
if cc := c.ContactCollector(); cc != nil {
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, senderName, "", "group")
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, senderName, "", "group", "user")
}
slog.Debug("zalo_personal group message recorded (no mention)",
@@ -144,7 +144,7 @@ func (c *Channel) handleGroupMessage(msg protocol.GroupMessage) {
// Collect contact for group-mentioned messages.
if cc := c.ContactCollector(); cc != nil {
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, senderName, "", "group")
cc.EnsureContact(ctx, c.Type(), c.Name(), senderID, senderID, senderName, "", "group", "user")
}
metadata := map[string]string{
@@ -105,10 +105,7 @@ func FetchGroups(ctx context.Context, sess *Session) ([]GroupListInfo, error) {
var allGroups []GroupListInfo
for i := 0; i < len(ids); i += batchSize {
end := i + batchSize
if end > len(ids) {
end = len(ids)
}
end := min(i+batchSize, len(ids))
batch := make(map[string]string, end-i)
for _, id := range ids[i:end] {
batch[id] = gridVerMap[id]
+7 -6
View File
@@ -19,11 +19,11 @@ func TestValidateSchedule(t *testing.T) {
sched Schedule
wantErr bool
}{
{"at_valid", Schedule{Kind: "at", AtMS: ptrInt64(time.Now().Add(time.Hour).UnixMilli())}, false},
{"at_valid", Schedule{Kind: "at", AtMS: new(time.Now().Add(time.Hour).UnixMilli())}, false},
{"at_missing_timestamp", Schedule{Kind: "at"}, true},
{"every_valid", Schedule{Kind: "every", EveryMS: ptrInt64(5000)}, false},
{"every_zero_interval", Schedule{Kind: "every", EveryMS: ptrInt64(0)}, true},
{"every_negative_interval", Schedule{Kind: "every", EveryMS: ptrInt64(-1)}, true},
{"every_valid", Schedule{Kind: "every", EveryMS: new(int64(5000))}, false},
{"every_zero_interval", Schedule{Kind: "every", EveryMS: new(int64(0))}, true},
{"every_negative_interval", Schedule{Kind: "every", EveryMS: new(int64(-1))}, true},
{"every_nil_interval", Schedule{Kind: "every"}, true},
{"cron_valid", Schedule{Kind: "cron", Expr: "*/5 * * * *"}, false},
{"cron_empty_expr", Schedule{Kind: "cron", Expr: ""}, true},
@@ -268,7 +268,7 @@ func TestService_NilHandler_NoPanic(t *testing.T) {
cs.Start()
time.Sleep(1500 * time.Millisecond) // wait for at least 1 tick
cs.Stop() // should not panic
cs.Stop() // should not panic
}
// --- Job failure with retry ---
@@ -354,4 +354,5 @@ func TestService_RunLog_PopulatedByAutoExecution(t *testing.T) {
// --- helpers ---
func ptrInt64(v int64) *int64 { return &v }
//go:fix inline
func ptrInt64(v int64) *int64 { return new(v) }
+9 -9
View File
@@ -82,13 +82,13 @@ func sweepExportTokens() {
// ExportManifest describes the archive contents.
type ExportManifest struct {
Version int `json:"version"`
Format string `json:"format"`
ExportedAt string `json:"exported_at"`
ExportedBy string `json:"exported_by"`
AgentKey string `json:"agent_key"`
AgentID string `json:"agent_id"`
Sections map[string]interface{} `json:"sections"`
Version int `json:"version"`
Format string `json:"format"`
ExportedAt string `json:"exported_at"`
ExportedBy string `json:"exported_by"`
AgentKey string `json:"agent_key"`
AgentID string `json:"agent_id"`
Sections map[string]any `json:"sections"`
}
// KGEntityExport is a portable KG entity (no internal UUID).
@@ -293,7 +293,7 @@ func (h *AgentsHandler) writeExportArchive(ctx context.Context, w io.Writer, ag
ExportedBy: store.UserIDFromContext(ctx),
AgentKey: ag.AgentKey,
AgentID: ag.ID.String(),
Sections: make(map[string]interface{}),
Sections: make(map[string]any),
}
// Section: config (always included)
@@ -693,7 +693,7 @@ func parseExportSections(raw string) map[string]bool {
return map[string]bool{"config": true, "context_files": true}
}
out := make(map[string]bool)
for _, s := range strings.Split(raw, ",") {
for s := range strings.SplitSeq(raw, ",") {
if s = strings.TrimSpace(s); s != "" {
out[s] = true
}
+1 -2
View File
@@ -96,7 +96,7 @@ func parseImportSections(raw string) map[string]bool {
return all
}
out := make(map[string]bool)
for _, s := range strings.Split(raw, ",") {
for s := range strings.SplitSeq(raw, ",") {
if s = strings.TrimSpace(s); s != "" {
out[s] = true
}
@@ -619,4 +619,3 @@ func (h *AgentsHandler) dedupAgentKey(ctx context.Context, base string) string {
}
return key
}
+5
View File
@@ -52,6 +52,11 @@ func (h *ChannelInstancesHandler) RegisterRoutes(mux *http.ServeMux) {
mux.HandleFunc("GET /v1/tenant-users", h.auth(h.handleListTenantUsers))
}
// Unified user search (contacts + tenant_users)
if h.contactStore != nil {
mux.HandleFunc("GET /v1/users/search", h.auth(h.handleSearchUsers))
}
// Group file writers (nested under channel instances)
if h.configPermStore != nil {
mux.HandleFunc("GET /v1/channels/instances/{id}/writers/groups", h.auth(h.handleWriterGroups))
+6
View File
@@ -37,6 +37,12 @@ func (h *SecureCLIHandler) RegisterRoutes(mux *http.ServeMux) {
mux.HandleFunc("PUT /v1/cli-credentials/{id}", h.auth(h.handleUpdate))
mux.HandleFunc("DELETE /v1/cli-credentials/{id}", h.auth(h.handleDelete))
mux.HandleFunc("POST /v1/cli-credentials/{id}/test", h.auth(h.handleDryRun))
// Per-user credential management
mux.HandleFunc("GET /v1/cli-credentials/{id}/user-credentials", h.auth(h.handleListUserCredentials))
mux.HandleFunc("GET /v1/cli-credentials/{id}/user-credentials/{userId}", h.auth(h.handleGetUserCredentials))
mux.HandleFunc("PUT /v1/cli-credentials/{id}/user-credentials/{userId}", h.auth(h.handleSetUserCredentials))
mux.HandleFunc("DELETE /v1/cli-credentials/{id}/user-credentials/{userId}", h.auth(h.handleDeleteUserCredentials))
}
func (h *SecureCLIHandler) auth(next http.HandlerFunc) http.HandlerFunc {
@@ -0,0 +1,147 @@
package http
import (
"encoding/json"
"net/http"
"github.com/google/uuid"
"github.com/nextlevelbuilder/goclaw/internal/i18n"
"github.com/nextlevelbuilder/goclaw/internal/store"
)
func (h *SecureCLIHandler) handleListUserCredentials(w http.ResponseWriter, r *http.Request) {
binaryID, err := uuid.Parse(r.PathValue("id"))
if err != nil {
locale := store.LocaleFromContext(r.Context())
writeJSON(w, http.StatusBadRequest, map[string]string{"error": i18n.T(locale, i18n.MsgInvalidID)})
return
}
creds, err := h.store.ListUserCredentials(r.Context(), binaryID)
if err != nil {
locale := store.LocaleFromContext(r.Context())
writeJSON(w, http.StatusInternalServerError, map[string]string{"error": i18n.T(locale, i18n.MsgInternalError, err.Error())})
return
}
// Return without encrypted env for listing (only user_id + timestamps)
type entry struct {
ID uuid.UUID `json:"id"`
BinaryID uuid.UUID `json:"binary_id"`
UserID string `json:"user_id"`
HasEnv bool `json:"has_env"`
CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"`
}
entries := make([]entry, 0, len(creds))
for _, c := range creds {
entries = append(entries, entry{
ID: c.ID,
BinaryID: c.BinaryID,
UserID: c.UserID,
HasEnv: len(c.EncryptedEnv) > 0,
CreatedAt: c.CreatedAt,
UpdatedAt: c.UpdatedAt,
})
}
writeJSON(w, http.StatusOK, map[string]any{"user_credentials": entries})
}
func (h *SecureCLIHandler) handleGetUserCredentials(w http.ResponseWriter, r *http.Request) {
binaryID, err := uuid.Parse(r.PathValue("id"))
if err != nil {
locale := store.LocaleFromContext(r.Context())
writeJSON(w, http.StatusBadRequest, map[string]string{"error": i18n.T(locale, i18n.MsgInvalidID)})
return
}
userID := r.PathValue("userId")
if userID == "" {
writeJSON(w, http.StatusBadRequest, map[string]string{"error": "user_id required"})
return
}
cred, err := h.store.GetUserCredentials(r.Context(), binaryID, userID)
if err != nil {
locale := store.LocaleFromContext(r.Context())
writeJSON(w, http.StatusInternalServerError, map[string]string{"error": i18n.T(locale, i18n.MsgInternalError, err.Error())})
return
}
if cred == nil {
writeJSON(w, http.StatusNotFound, map[string]string{"error": "not found"})
return
}
// Return decrypted env as JSON object (admin-only endpoint)
var envObj any
if len(cred.EncryptedEnv) > 0 {
_ = json.Unmarshal(cred.EncryptedEnv, &envObj)
}
writeJSON(w, http.StatusOK, map[string]any{
"user_id": cred.UserID,
"env": envObj,
})
}
func (h *SecureCLIHandler) handleSetUserCredentials(w http.ResponseWriter, r *http.Request) {
binaryID, err := uuid.Parse(r.PathValue("id"))
if err != nil {
locale := store.LocaleFromContext(r.Context())
writeJSON(w, http.StatusBadRequest, map[string]string{"error": i18n.T(locale, i18n.MsgInvalidID)})
return
}
userID := r.PathValue("userId")
if userID == "" {
writeJSON(w, http.StatusBadRequest, map[string]string{"error": "user_id required"})
return
}
var body struct {
Env json.RawMessage `json:"env"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid JSON body"})
return
}
if len(body.Env) == 0 {
writeJSON(w, http.StatusBadRequest, map[string]string{"error": "env is required"})
return
}
// Validate env is a JSON object
var envCheck map[string]string
if err := json.Unmarshal(body.Env, &envCheck); err != nil {
writeJSON(w, http.StatusBadRequest, map[string]string{"error": "env must be a JSON object with string values"})
return
}
if err := h.store.SetUserCredentials(r.Context(), binaryID, userID, body.Env); err != nil {
locale := store.LocaleFromContext(r.Context())
writeJSON(w, http.StatusInternalServerError, map[string]string{"error": i18n.T(locale, i18n.MsgInternalError, err.Error())})
return
}
h.emitCacheInvalidate("")
writeJSON(w, http.StatusOK, map[string]string{"status": "ok"})
}
func (h *SecureCLIHandler) handleDeleteUserCredentials(w http.ResponseWriter, r *http.Request) {
binaryID, err := uuid.Parse(r.PathValue("id"))
if err != nil {
locale := store.LocaleFromContext(r.Context())
writeJSON(w, http.StatusBadRequest, map[string]string{"error": i18n.T(locale, i18n.MsgInvalidID)})
return
}
userID := r.PathValue("userId")
if userID == "" {
writeJSON(w, http.StatusBadRequest, map[string]string{"error": "user_id required"})
return
}
if err := h.store.DeleteUserCredentials(r.Context(), binaryID, userID); err != nil {
locale := store.LocaleFromContext(r.Context())
writeJSON(w, http.StatusInternalServerError, map[string]string{"error": i18n.T(locale, i18n.MsgInternalError, err.Error())})
return
}
h.emitCacheInvalidate("")
writeJSON(w, http.StatusOK, map[string]string{"status": "ok"})
}
+4 -5
View File
@@ -94,12 +94,12 @@ func (h *SkillsHandler) doSkillsImport(ctx context.Context, r io.Reader, userID
continue
}
rest := strings.TrimPrefix(name, "skills/")
slashIdx := strings.Index(rest, "/")
if slashIdx < 0 {
before, after, ok := strings.Cut(rest, "/")
if !ok {
continue
}
slug := rest[:slashIdx]
file := rest[slashIdx+1:]
slug := before
file := after
if slug == "" || file == "" {
continue
}
@@ -240,4 +240,3 @@ func tenantIDForSkillImport(ctx context.Context) uuid.UUID {
}
return tid
}
+1 -1
View File
@@ -18,7 +18,7 @@ type ProgressEvent struct {
// sendSSE writes a named SSE event with JSON payload and flushes.
// The event format follows the standard SSE spec: "event: <name>\ndata: <json>\n\n".
func sendSSE(w http.ResponseWriter, flusher http.Flusher, event string, data interface{}) {
func sendSSE(w http.ResponseWriter, flusher http.Flusher, event string, data any) {
jsonData, err := json.Marshal(data)
if err != nil {
return
+22 -22
View File
@@ -24,14 +24,14 @@ import (
// TeamExportManifest describes the contents of a team export archive.
type TeamExportManifest struct {
Version int `json:"version"`
Format string `json:"format"`
ExportedAt string `json:"exported_at"`
ExportedBy string `json:"exported_by"`
TeamName string `json:"team_name"`
TeamID string `json:"team_id"`
AgentKeys []string `json:"agent_keys"`
Sections map[string]interface{} `json:"sections"`
Version int `json:"version"`
Format string `json:"format"`
ExportedAt string `json:"exported_at"`
ExportedBy string `json:"exported_by"`
TeamName string `json:"team_name"`
TeamID string `json:"team_id"`
AgentKeys []string `json:"agent_keys"`
Sections map[string]any `json:"sections"`
}
// handleTeamExportPreview returns team export counts without building the archive.
@@ -60,12 +60,12 @@ func (h *AgentsHandler) handleTeamExportPreview(w http.ResponseWriter, r *http.R
}
writeJSON(w, http.StatusOK, map[string]any{
"team_name": teamMeta.Name,
"team_id": teamIDStr,
"tasks": tasks,
"members": members,
"agent_links": links,
"agent_count": len(agentMembers),
"team_name": teamMeta.Name,
"team_id": teamIDStr,
"tasks": tasks,
"members": members,
"agent_links": links,
"agent_count": len(agentMembers),
})
}
@@ -148,7 +148,7 @@ func (h *AgentsHandler) writeTeamExportArchive(ctx context.Context, w io.Writer,
TeamName: teamMeta.Name,
TeamID: teamID.String(),
AgentKeys: []string{},
Sections: make(map[string]interface{}),
Sections: make(map[string]any),
}
// team/team.json
@@ -273,13 +273,13 @@ func (h *AgentsHandler) writeTeamExportArchive(ctx context.Context, w io.Writer,
}
sections := map[string]bool{
"context_files": true,
"memory": true,
"context_files": true,
"memory": true,
"knowledge_graph": true,
"cron": true,
"user_profiles": true,
"user_overrides": true,
"workspace": true,
"cron": true,
"user_profiles": true,
"user_overrides": true,
"workspace": true,
}
for _, member := range agentMembers {
@@ -467,7 +467,7 @@ func (h *AgentsHandler) writeAgentSectionsToTar(ctx context.Context, tw *tar.Wri
exportRelations = append(exportRelations, KGRelationExport{
SourceExternalID: idToExternal[rel.SourceEntityID],
TargetExternalID: idToExternal[rel.TargetEntityID],
UserID: rel.UserID, RelationType: rel.RelationType,
UserID: rel.UserID, RelationType: rel.RelationType,
Confidence: rel.Confidence, Properties: rel.Properties,
})
}
+10 -10
View File
@@ -174,12 +174,12 @@ func readTeamImportArchive(r io.Reader) (*teamImportArchive, error) {
continue
}
rest := strings.TrimPrefix(name, "agents/")
slashIdx := strings.Index(rest, "/")
if slashIdx < 0 {
before, after, ok := strings.Cut(rest, "/")
if !ok {
continue
}
agentKey := rest[:slashIdx]
relPath := rest[slashIdx+1:]
agentKey := before
relPath := after
if agentKey == "" || relPath == "" {
continue
}
@@ -358,13 +358,13 @@ func (h *AgentsHandler) doTeamImport(ctx context.Context, r *http.Request, teamA
}
sections := map[string]bool{
"context_files": true,
"memory": true,
"context_files": true,
"memory": true,
"knowledge_graph": true,
"cron": true,
"user_profiles": true,
"user_overrides": true,
"workspace": true,
"cron": true,
"user_profiles": true,
"user_overrides": true,
"workspace": true,
}
if _, err := h.doMergeImport(ctx, ag, agArc, sections, progressFn); err != nil {
slog.Warn("team.import: merge agent data failed", "key", dedupedKey, "error", err)
+116
View File
@@ -0,0 +1,116 @@
package http
import (
"log/slog"
"net/http"
"strconv"
"strings"
"github.com/google/uuid"
"github.com/nextlevelbuilder/goclaw/internal/store"
)
// UserSearchResult is a unified result from contacts + tenant_users.
type UserSearchResult struct {
ID string `json:"id"`
DisplayName *string `json:"display_name,omitempty"`
Username *string `json:"username,omitempty"`
Source string `json:"source"` // "contact" or "tenant_user"
ChannelType *string `json:"channel_type,omitempty"`
PeerKind *string `json:"peer_kind,omitempty"`
MergedTenantUserID *string `json:"merged_tenant_user_id,omitempty"`
Role *string `json:"role,omitempty"`
}
// handleSearchUsers returns unified results from channel_contacts + tenant_users.
// GET /v1/users/search?q=&limit=30&peer_kind=
// Empty q → return most recent. With q → ILIKE search across both tables.
func (h *ChannelInstancesHandler) handleSearchUsers(w http.ResponseWriter, r *http.Request) {
q := r.URL.Query().Get("q")
peerKind := r.URL.Query().Get("peer_kind")
source := r.URL.Query().Get("source") // "contact", "tenant_user", or "" (both)
limit := 30
if v := r.URL.Query().Get("limit"); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 && n <= 100 {
limit = n
}
}
ctx := r.Context()
tid := store.TenantIDFromContext(ctx)
var results []UserSearchResult
mergedUserIDs := make(map[string]bool) // for deduplication between contacts and tenant_users
// 1. Search channel_contacts (skip if source=tenant_user)
if h.contactStore != nil && source != "tenant_user" {
opts := store.ContactListOpts{
Search: q,
PeerKind: peerKind,
Limit: limit,
}
contacts, err := h.contactStore.ListContacts(ctx, opts)
if err != nil {
slog.Warn("user_search.contacts", "error", err)
}
for _, c := range contacts {
r := UserSearchResult{
ID: c.SenderID,
DisplayName: c.DisplayName,
Username: c.Username,
Source: "contact",
ChannelType: &c.ChannelType,
PeerKind: c.PeerKind,
}
if c.MergedID != nil {
if resolved, err := h.contactStore.ResolveTenantUserID(ctx, c.ChannelType, c.SenderID); err == nil && resolved != "" {
r.MergedTenantUserID = &resolved
mergedUserIDs[resolved] = true
}
}
results = append(results, r)
}
}
// 2. Search tenant_users (skip if source=contact)
if h.tenantStore != nil && tid != uuid.Nil && source != "contact" {
users, err := h.tenantStore.ListUsers(ctx, tid)
if err != nil {
slog.Warn("user_search.tenant_users", "error", err)
}
for _, u := range users {
if mergedUserIDs[u.UserID] {
continue
}
if q != "" && !containsInsensitive(u.UserID, q) && !containsInsensitive(ptrStr(u.DisplayName), q) {
continue
}
if len(results) >= limit {
break
}
role := u.Role
results = append(results, UserSearchResult{
ID: u.UserID,
DisplayName: u.DisplayName,
Source: "tenant_user",
Role: &role,
})
}
}
if results == nil {
results = []UserSearchResult{}
}
writeJSON(w, http.StatusOK, map[string]any{"results": results})
}
func ptrStr(s *string) string {
if s == nil {
return ""
}
return *s
}
func containsInsensitive(s, substr string) bool {
return strings.Contains(strings.ToLower(s), strings.ToLower(substr))
}
+1 -4
View File
@@ -16,10 +16,7 @@ func JaroWinkler(a, b string) float64 {
}
// Jaro similarity
matchDist := max(len(a), len(b))/2 - 1
if matchDist < 0 {
matchDist = 0
}
matchDist := max(max(len(a), len(b))/2-1, 0)
aMatched := make([]bool, len(a))
bMatched := make([]bool, len(b))
+3
View File
@@ -90,6 +90,9 @@ func (t *BridgeTool) ServerName() string { return t.serverName }
// OriginalName returns the original MCP tool name (without prefix).
func (t *BridgeTool) OriginalName() string { return t.toolName }
// IsConnected returns whether the underlying MCP server connection is healthy.
func (t *BridgeTool) IsConnected() bool { return t.connected.Load() }
func (t *BridgeTool) Execute(ctx context.Context, args map[string]any) *tools.Result {
if !t.connected.Load() {
return tools.ErrorResult(fmt.Sprintf("MCP server %q is disconnected", t.serverName))
+14
View File
@@ -83,6 +83,11 @@ type Manager struct {
deferredTools map[string]*BridgeTool // registeredName → BridgeTool
activatedTools map[string]struct{} // tracks activated tool names for group:mcp
searchMode bool
// User-credential servers: servers requiring per-user credentials, stored during
// LoadForAgent("") for later per-request tool resolution. These servers are NOT
// connected at startup — connections are created per-user via pool.AcquireUser().
userCredServers []store.MCPAccessInfo
}
// ManagerOption configures the Manager.
@@ -276,8 +281,17 @@ func (m *Manager) LoadForAgent(ctx context.Context, agentID uuid.UUID, userID st
// Unregister all existing MCP tools first
m.unregisterAllTools()
m.userCredServers = nil
for _, info := range accessible {
// When loading at startup (userID=""), store servers requiring per-user
// credentials for later per-request resolution instead of skipping them.
if userID == "" && requireUserCreds(info.Server.Settings) && info.Server.Enabled {
m.userCredServers = append(m.userCredServers, info)
slog.Debug("mcp.server.deferred_user_creds", "server", info.Server.Name)
continue
}
rs := m.resolveServerCredentials(ctx, info, userID)
if rs == nil {
continue
+8
View File
@@ -7,9 +7,17 @@ import (
"time"
mcpgo "github.com/mark3labs/mcp-go/mcp"
"github.com/nextlevelbuilder/goclaw/internal/store"
"github.com/nextlevelbuilder/goclaw/internal/tools"
)
// UserCredServers returns servers requiring per-user credentials.
// These are stored during LoadForAgent("") and used by the agent loop
// for per-request tool resolution via pool.AcquireUser().
func (m *Manager) UserCredServers() []store.MCPAccessInfo {
return m.userCredServers
}
// ToolNames returns all registered MCP tool names.
func (m *Manager) ToolNames() []string {
m.mu.RLock()
+6
View File
@@ -35,6 +35,12 @@ func joinErrors(errs []string) string {
return result.String()
}
// ParseJSONBytesToStringSlice converts JSONB []byte to []string (exported for loop_mcp_user).
func ParseJSONBytesToStringSlice(data []byte) []string { return jsonBytesToStringSlice(data) }
// ParseJSONBytesToStringMap converts JSONB []byte to map[string]string (exported for loop_mcp_user).
func ParseJSONBytesToStringMap(data []byte) map[string]string { return jsonBytesToStringMap(data) }
// jsonBytesToStringSlice converts JSONB []byte to []string. Returns nil on error.
func jsonBytesToStringSlice(data []byte) []string {
if len(data) == 0 {
+293 -49
View File
@@ -4,28 +4,37 @@ import (
"context"
"fmt"
"log/slog"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/google/uuid"
mcpclient "github.com/mark3labs/mcp-go/client"
mcpgo "github.com/mark3labs/mcp-go/mcp"
)
// PoolConfig configures the MCP connection pool.
type PoolConfig struct {
MaxSize int // global max connections (default 100)
MaxIdle int // max idle connections to keep alive (default 20)
IdleTTL time.Duration // close idle connections after this (default 20m)
AcquireTimeout time.Duration // wait for pool slot before error (default 60s)
MaxSize int // global max connections (default 200)
MaxIdle int // max idle connections to keep alive (default 20)
IdleTTL time.Duration // close idle connections after this (default 20m)
AcquireTimeout time.Duration // wait for pool slot before error (default 60s)
MaxUserConns int // max per-user connections per MCP server (default 30)
UserIdleTTL time.Duration // close idle user connections after this (default 15m)
UserAcquireTimeout time.Duration // wait for user pool slot before error (default 10s)
}
// DefaultPoolConfig returns the default pool configuration.
func DefaultPoolConfig() PoolConfig {
return PoolConfig{
MaxSize: 100,
MaxIdle: 20,
IdleTTL: 20 * time.Minute,
AcquireTimeout: 60 * time.Second,
MaxSize: 200,
MaxIdle: 20,
IdleTTL: 20 * time.Minute,
AcquireTimeout: 60 * time.Second,
MaxUserConns: 30,
UserIdleTTL: 15 * time.Minute,
UserAcquireTimeout: 10 * time.Second,
}
}
@@ -39,18 +48,21 @@ type poolEntry struct {
// Pool manages shared MCP server connections across agents.
// Connections are keyed by tenantID/serverName for tenant isolation.
// Per-user connections are keyed by tenantID/serverName/user:userID.
type Pool struct {
mu sync.Mutex
servers map[string]*poolEntry
cfg PoolConfig
slot chan struct{} // semaphore for MaxSize
stopCh chan struct{}
mu sync.Mutex
servers map[string]*poolEntry // shared connections: tenantID/serverName
userServers map[string]*poolEntry // user connections: tenantID/serverName/user:userID
userSlots map[string]chan struct{} // per-server semaphores: tenantID/serverName → capacity MaxUserConns
cfg PoolConfig
slot chan struct{} // semaphore for MaxSize
stopCh chan struct{}
}
// NewPool creates a shared MCP connection pool with idle eviction.
func NewPool(cfg PoolConfig) *Pool {
if cfg.MaxSize <= 0 {
cfg.MaxSize = 100
cfg.MaxSize = 200
}
if cfg.MaxIdle <= 0 {
cfg.MaxIdle = 20
@@ -61,12 +73,23 @@ func NewPool(cfg PoolConfig) *Pool {
if cfg.AcquireTimeout <= 0 {
cfg.AcquireTimeout = 60 * time.Second
}
if cfg.MaxUserConns <= 0 {
cfg.MaxUserConns = 30
}
if cfg.UserIdleTTL <= 0 {
cfg.UserIdleTTL = 15 * time.Minute
}
if cfg.UserAcquireTimeout <= 0 {
cfg.UserAcquireTimeout = 10 * time.Second
}
p := &Pool{
servers: make(map[string]*poolEntry),
cfg: cfg,
slot: make(chan struct{}, cfg.MaxSize),
stopCh: make(chan struct{}),
servers: make(map[string]*poolEntry),
userServers: make(map[string]*poolEntry),
userSlots: make(map[string]chan struct{}),
cfg: cfg,
slot: make(chan struct{}, cfg.MaxSize),
stopCh: make(chan struct{}),
}
go p.evictLoop()
return p
@@ -77,6 +100,17 @@ func poolKey(tenantID uuid.UUID, name string) string {
return tenantID.String() + "/" + name
}
// UserPoolKey builds a tenant+user-scoped key for user pool lookups.
// Exported for callers that need to construct release keys.
func UserPoolKey(tenantID uuid.UUID, serverName, userID string) string {
return tenantID.String() + "/" + serverName + "/user:" + userID
}
// userSlotKey returns the per-server semaphore key (tenantID/serverName).
func userSlotKey(tenantID uuid.UUID, serverName string) string {
return tenantID.String() + "/" + serverName
}
// Acquire returns a shared connection for the named server scoped to a tenant.
// If no connection exists, it connects using the provided config.
// Blocks up to AcquireTimeout if pool is at MaxSize.
@@ -161,6 +195,99 @@ func (p *Pool) Acquire(ctx context.Context, tenantID uuid.UUID, name, transportT
return entry, nil
}
// AcquireUser returns a per-user connection for the named server scoped to a tenant+user.
// If no connection exists, it connects using the provided config.
// Blocks up to UserAcquireTimeout if per-server user slot limit is reached.
func (p *Pool) AcquireUser(ctx context.Context, tenantID uuid.UUID, name, userID, transportType, command string, args []string, env map[string]string, url string, headers map[string]string, timeoutSec int) (*poolEntry, error) {
key := UserPoolKey(tenantID, name, userID)
slotKey := userSlotKey(tenantID, name)
p.mu.Lock()
if entry, ok := p.userServers[key]; ok && entry.state.connected.Load() {
entry.refCount++
entry.lastUsed = time.Now()
p.mu.Unlock()
slog.Debug("mcp.pool.user.reuse", "key", key, "refCount", entry.refCount)
return entry, nil
}
// If entry exists but disconnected, close old and reclaim slot
if old, ok := p.userServers[key]; ok {
if old.state.cancel != nil {
old.state.cancel()
}
if old.state.client != nil {
_ = old.state.client.Close()
}
delete(p.userServers, key)
// Return slot to per-server semaphore
if sem, ok := p.userSlots[slotKey]; ok {
select {
case <-sem:
default:
}
}
}
// Ensure per-server semaphore exists (lazy init under lock)
if _, ok := p.userSlots[slotKey]; !ok {
p.userSlots[slotKey] = make(chan struct{}, p.cfg.MaxUserConns)
}
sem := p.userSlots[slotKey]
p.mu.Unlock()
// Acquire a user slot for this server (blocks up to UserAcquireTimeout)
if err := p.acquireUserSlot(ctx, sem, slotKey); err != nil {
return nil, fmt.Errorf("mcp user pool exhausted for server %s: %w", name, err)
}
// Connect outside the lock (may be slow)
ss, mcpTools, err := connectAndDiscover(ctx, name, transportType, command, args, env, url, headers, timeoutSec)
if err != nil {
// Return slot on failure
select {
case <-sem:
default:
}
return nil, err
}
// Start health loop
hctx, hcancel := context.WithCancel(context.Background())
ss.cancel = hcancel
go poolHealthLoop(hctx, ss)
entry := &poolEntry{
state: ss,
tools: mcpTools,
refCount: 1,
lastUsed: time.Now(),
}
p.mu.Lock()
// Check if another goroutine connected while we were connecting
if existing, ok := p.userServers[key]; ok && existing.state.connected.Load() {
p.mu.Unlock()
hcancel()
_ = ss.client.Close()
// Return our extra slot
select {
case <-sem:
default:
}
p.mu.Lock()
existing.refCount++
existing.lastUsed = time.Now()
p.mu.Unlock()
return existing, nil
}
p.userServers[key] = entry
p.mu.Unlock()
slog.Info("mcp.pool.user.connected", "key", key, "tools", len(mcpTools))
return entry, nil
}
// acquireSlot tries to acquire a pool slot, evicting idle connections if needed.
func (p *Pool) acquireSlot(ctx context.Context) error {
// Fast path: slot available
@@ -197,6 +324,29 @@ func (p *Pool) acquireSlot(ctx context.Context) error {
}
}
// acquireUserSlot tries to acquire a per-server user slot.
func (p *Pool) acquireUserSlot(ctx context.Context, sem chan struct{}, slotKey string) error {
// Fast path: slot available
select {
case sem <- struct{}{}:
return nil
default:
}
// Wait up to UserAcquireTimeout
timer := time.NewTimer(p.cfg.UserAcquireTimeout)
defer timer.Stop()
select {
case sem <- struct{}{}:
return nil
case <-timer.C:
return fmt.Errorf("timeout after %s waiting for user slot (max %d, server %s)", p.cfg.UserAcquireTimeout, p.cfg.MaxUserConns, slotKey)
case <-ctx.Done():
return ctx.Err()
}
}
// Release decrements the reference count for a server.
// Accepts the same key format as Acquire (tenantID + name).
func (p *Pool) Release(key string) {
@@ -213,6 +363,22 @@ func (p *Pool) Release(key string) {
}
}
// ReleaseUser decrements the reference count for a user-scoped connection.
// Accepts the same key format as AcquireUser (tenantID + serverName + userID).
func (p *Pool) ReleaseUser(key string) {
p.mu.Lock()
defer p.mu.Unlock()
if entry, ok := p.userServers[key]; ok {
entry.refCount--
if entry.refCount < 0 {
entry.refCount = 0
}
entry.lastUsed = time.Now()
slog.Debug("mcp.pool.user.release", "key", key, "refCount", entry.refCount)
}
}
// Stop closes all pooled connections and stops eviction. Called on gateway shutdown.
func (p *Pool) Stop() {
close(p.stopCh)
@@ -230,6 +396,17 @@ func (p *Pool) Stop() {
slog.Debug("mcp.pool.stopped", "key", key)
}
p.servers = make(map[string]*poolEntry)
for key, entry := range p.userServers {
if entry.state.cancel != nil {
entry.state.cancel()
}
if entry.state.client != nil {
_ = entry.state.client.Close()
}
slog.Debug("mcp.pool.user.stopped", "key", key)
}
p.userServers = make(map[string]*poolEntry)
}
// Evict closes a specific pooled connection by tenant + server name.
@@ -273,11 +450,14 @@ func (p *Pool) evictLoop() {
}
// evictIdle closes connections idle > IdleTTL when total idle exceeds MaxIdle.
// Also evicts user connections idle > UserIdleTTL.
func (p *Pool) evictIdle() {
p.mu.Lock()
defer p.mu.Unlock()
now := time.Now()
// Evict shared connections
var idleKeys []string
for key, entry := range p.servers {
if entry.refCount == 0 && now.Sub(entry.lastUsed) > p.cfg.IdleTTL {
@@ -295,39 +475,74 @@ func (p *Pool) evictIdle() {
// Only evict if over MaxIdle
toEvict := totalIdle - p.cfg.MaxIdle
if toEvict <= 0 && len(idleKeys) == 0 {
return
if toEvict > 0 || len(idleKeys) > 0 {
for _, key := range idleKeys {
entry := p.servers[key]
if entry.state.cancel != nil {
entry.state.cancel()
}
if entry.state.client != nil {
_ = entry.state.client.Close()
}
delete(p.servers, key)
select {
case <-p.slot:
default:
}
slog.Debug("mcp.pool.evicted", "key", key, "reason", "idle_ttl")
}
}
// Evict TTL-expired first, then oldest if still over MaxIdle
for _, key := range idleKeys {
entry := p.servers[key]
if entry.state.cancel != nil {
entry.state.cancel()
// Evict user connections idle > UserIdleTTL
for key, entry := range p.userServers {
if entry.refCount == 0 && now.Sub(entry.lastUsed) > p.cfg.UserIdleTTL {
if entry.state.cancel != nil {
entry.state.cancel()
}
if entry.state.client != nil {
_ = entry.state.client.Close()
}
delete(p.userServers, key)
// Return slot to per-server semaphore
// Extract slotKey from user key: "tenantID/serverName/user:userID" → "tenantID/serverName"
// We search userSlots by iterating — key format guarantees prefix match
for slotKey, sem := range p.userSlots {
if strings.HasPrefix(key, slotKey+"/") {
select {
case <-sem:
default:
}
break
}
}
slog.Debug("mcp.pool.user.evicted", "key", key, "reason", "idle_ttl")
}
if entry.state.client != nil {
_ = entry.state.client.Close()
}
delete(p.servers, key)
// Return slot
select {
case <-p.slot:
default:
}
slog.Debug("mcp.pool.evicted", "key", key, "reason", "idle_ttl")
}
}
// evictOldestIdleLocked evicts one idle entry (oldest lastUsed). Caller must hold mu.
// evictOldestIdleLocked evicts one idle entry (oldest lastUsed) from shared or user pools.
// Caller must hold mu.
func (p *Pool) evictOldestIdleLocked() bool {
var oldestKey string
var oldestTime time.Time
isUser := false
for key, entry := range p.servers {
if entry.refCount == 0 {
if oldestKey == "" || entry.lastUsed.Before(oldestTime) {
oldestKey = key
oldestTime = entry.lastUsed
isUser = false
}
}
}
for key, entry := range p.userServers {
if entry.refCount == 0 {
if oldestKey == "" || entry.lastUsed.Before(oldestTime) {
oldestKey = key
oldestTime = entry.lastUsed
isUser = true
}
}
}
@@ -336,23 +551,52 @@ func (p *Pool) evictOldestIdleLocked() bool {
return false
}
entry := p.servers[oldestKey]
if entry.state.cancel != nil {
entry.state.cancel()
if isUser {
entry := p.userServers[oldestKey]
if entry.state.cancel != nil {
entry.state.cancel()
}
if entry.state.client != nil {
_ = entry.state.client.Close()
}
delete(p.userServers, oldestKey)
for slotKey, sem := range p.userSlots {
if strings.HasPrefix(oldestKey, slotKey+"/") {
select {
case <-sem:
default:
}
break
}
}
} else {
entry := p.servers[oldestKey]
if entry.state.cancel != nil {
entry.state.cancel()
}
if entry.state.client != nil {
_ = entry.state.client.Close()
}
delete(p.servers, oldestKey)
select {
case <-p.slot:
default:
}
}
if entry.state.client != nil {
_ = entry.state.client.Close()
}
delete(p.servers, oldestKey)
// Return slot to semaphore
select {
case <-p.slot:
default:
}
slog.Debug("mcp.pool.evicted", "key", oldestKey, "reason", "make_room")
slog.Debug("mcp.pool.evicted", "key", oldestKey, "reason", "make_room", "user", isUser)
return true
}
// Client returns the MCP client for this pool entry.
func (e *poolEntry) Client() *mcpclient.Client { return e.state.client }
// Connected returns a pointer to the connected flag for this pool entry.
func (e *poolEntry) Connected() *atomic.Bool { return &e.state.connected }
// MCPTools returns the discovered MCP tool definitions for this pool entry.
func (e *poolEntry) MCPTools() []mcpgo.Tool { return e.tools }
// poolHealthLoop is a standalone health loop for pool-managed connections.
func poolHealthLoop(ctx context.Context, ss *serverState) {
ticker := newHealthTicker()
+4 -4
View File
@@ -107,7 +107,7 @@ func TestComputeDelay_JitterRange(t *testing.T) {
min := 750 * time.Millisecond
max := 1250 * time.Millisecond
for i := 0; i < 100; i++ {
for range 100 {
d := computeDelay(cfg, 1, err)
if d < min || d > max {
t.Fatalf("jitter out of range: got %v, want [%v, %v]", d, min, max)
@@ -123,7 +123,7 @@ func TestComputeDelay_NeverNegative(t *testing.T) {
}
err := &HTTPError{Status: 500}
for i := 0; i < 200; i++ {
for range 200 {
d := computeDelay(cfg, 1, err)
if d < 0 {
t.Fatalf("negative delay: %v", d)
@@ -158,8 +158,8 @@ func TestParseRetryAfter(t *testing.T) {
{"integer_seconds", "30", 30 * time.Second},
{"zero", "0", 0},
{"negative_int", "-5", -5 * time.Second}, // strconv.Atoi succeeds → returns negative duration (caller should clamp)
{"non_numeric", "abc", 0}, // neither int nor date
{"float", "1.5", 0}, // not a valid int
{"non_numeric", "abc", 0}, // neither int nor date
{"float", "1.5", 0}, // not a valid int
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
+2 -2
View File
@@ -270,7 +270,7 @@ func TestLane_ConcurrencyEnforcement(t *testing.T) {
var current atomic.Int32
var wg sync.WaitGroup
for i := 0; i < 10; i++ {
for range 10 {
wg.Add(1)
err := lane.Submit(context.Background(), func() {
defer wg.Done()
@@ -400,7 +400,7 @@ func TestSessionQueue_Debounce_CollapsesRapidMessages(t *testing.T) {
// Send 5 rapid messages within debounce window
var channels []<-chan RunOutcome
for i := 0; i < 5; i++ {
for i := range 5 {
ch := sq.Enqueue(ctx, agent.RunRequest{
RunID: "r" + string(rune('0'+i)),
SessionKey: "test",
+2 -2
View File
@@ -86,7 +86,7 @@ func TestAddMessage_ConcurrentSafety(t *testing.T) {
const n = 100
var wg sync.WaitGroup
wg.Add(n)
for i := 0; i < n; i++ {
for i := range n {
go func(i int) {
defer wg.Done()
m.AddMessage(ctx, key, providers.Message{Role: "user", Content: "msg"})
@@ -209,7 +209,7 @@ func TestTruncateHistory(t *testing.T) {
ctx := context.Background()
key := "agent:a1:s1"
for i := 0; i < 10; i++ {
for range 10 {
m.AddMessage(ctx, key, providers.Message{Role: "user", Content: "msg"})
}
+11 -2
View File
@@ -23,14 +23,23 @@ func NewContactCollector(s ContactStore, c cache.Cache[bool]) *ContactCollector
}
// EnsureContact creates or refreshes a contact entry, skipping DB if recently seen.
func (c *ContactCollector) EnsureContact(ctx context.Context, channelType, channelInstance, senderID, userID, displayName, username, peerKind string) {
// contactType: "user" (individual sender) or "group" (group chat entity).
func (c *ContactCollector) EnsureContact(ctx context.Context, channelType, channelInstance, senderID, userID, displayName, username, peerKind, contactType string) {
key := channelType + ":" + senderID
if _, ok := c.seen.Get(ctx, key); ok {
return
}
if err := c.store.UpsertContact(ctx, channelType, channelInstance, senderID, userID, displayName, username, peerKind); err != nil {
if contactType == "" {
contactType = "user"
}
if err := c.store.UpsertContact(ctx, channelType, channelInstance, senderID, userID, displayName, username, peerKind, contactType); err != nil {
slog.Warn("contact_collector.upsert_failed", "error", err, "channel", channelType, "sender", senderID)
return
}
c.seen.Set(ctx, key, true, contactSeenTTL)
}
// ResolveTenantUserID delegates to the underlying ContactStore.
func (c *ContactCollector) ResolveTenantUserID(ctx context.Context, channelType, senderID string) (string, error) {
return c.store.ResolveTenantUserID(ctx, channelType, senderID)
}
+7 -1
View File
@@ -19,6 +19,7 @@ type ChannelContact struct {
Username *string `json:"username,omitempty"`
AvatarURL *string `json:"avatar_url,omitempty"`
PeerKind *string `json:"peer_kind,omitempty"`
ContactType string `json:"contact_type"` // "user" or "group"
MergedID *uuid.UUID `json:"merged_id,omitempty"`
FirstSeenAt time.Time `json:"first_seen_at"`
LastSeenAt time.Time `json:"last_seen_at"`
@@ -37,7 +38,7 @@ type ContactListOpts struct {
type ContactStore interface {
// UpsertContact creates or updates a contact. On conflict (channel_type, sender_id),
// updates display_name, username, user_id, channel_instance, and last_seen_at.
UpsertContact(ctx context.Context, channelType, channelInstance, senderID, userID, displayName, username, peerKind string) error
UpsertContact(ctx context.Context, channelType, channelInstance, senderID, userID, displayName, username, peerKind, contactType string) error
// ListContacts searches contacts with pagination and filters.
ListContacts(ctx context.Context, opts ContactListOpts) ([]ChannelContact, error)
@@ -63,4 +64,9 @@ type ContactStore interface {
// GetContactsByMergedID returns all contacts linked to a given merged_id.
// Tenant-scoped via context.
GetContactsByMergedID(ctx context.Context, mergedID uuid.UUID) ([]ChannelContact, error)
// ResolveTenantUserID looks up a contact by (channelType, senderID) and, if
// the contact has been merged, returns the linked tenant_user's user_id.
// Returns ("", nil) when the contact is not found or not merged.
ResolveTenantUserID(ctx context.Context, channelType, senderID string) (string, error)
}
+8 -6
View File
@@ -10,10 +10,10 @@ func TestNextRunForToggle_DisableClearsNextRun(t *testing.T) {
now := time.Date(2026, time.March, 28, 12, 0, 0, 0, time.UTC)
schedule := &CronSchedule{
Kind: "every",
EveryMS: int64Ptr(60_000),
EveryMS: new(int64(60_000)),
}
next, err := NextRunForToggle(schedule, false, true, timePtr(now.Add(time.Minute)), now, "")
next, err := NextRunForToggle(schedule, false, true, new(now.Add(time.Minute)), now, "")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
@@ -26,7 +26,7 @@ func TestNextRunForToggle_EnableRecomputesEverySchedule(t *testing.T) {
now := time.Date(2026, time.March, 28, 12, 0, 0, 0, time.UTC)
schedule := &CronSchedule{
Kind: "every",
EveryMS: int64Ptr(60_000),
EveryMS: new(int64(60_000)),
}
next, err := NextRunForToggle(schedule, true, false, nil, now, "")
@@ -69,7 +69,7 @@ func TestNextRunForToggle_AlreadyEnabledPreservesCurrentNextRun(t *testing.T) {
currentNextRun := now.Add(5 * time.Minute)
schedule := &CronSchedule{
Kind: "every",
EveryMS: int64Ptr(60_000),
EveryMS: new(int64(60_000)),
}
next, err := NextRunForToggle(schedule, true, true, &currentNextRun, now.Add(time.Minute), "")
@@ -104,10 +104,12 @@ func TestNextRunForToggle_ExpiredAtReturnsError(t *testing.T) {
}
}
//go:fix inline
func int64Ptr(v int64) *int64 {
return &v
return new(v)
}
//go:fix inline
func timePtr(v time.Time) *time.Time {
return &v
return new(v)
}
+20 -12
View File
@@ -14,30 +14,32 @@ import (
// PGContactStore implements store.ContactStore backed by Postgres.
type PGContactStore struct {
db *sql.DB
db *sql.DB
resolveCache *contactResolveCache // tenant-user resolution cache (60s TTL)
}
// NewPGContactStore creates a new PGContactStore.
func NewPGContactStore(db *sql.DB) *PGContactStore {
return &PGContactStore{db: db}
return &PGContactStore{db: db, resolveCache: newContactResolveCache()}
}
func (s *PGContactStore) UpsertContact(ctx context.Context, channelType, channelInstance, senderID, userID, displayName, username, peerKind string) error {
func (s *PGContactStore) UpsertContact(ctx context.Context, channelType, channelInstance, senderID, userID, displayName, username, peerKind, contactType string) error {
tenantID := store.TenantIDFromContext(ctx)
if tenantID == uuid.Nil {
tenantID = store.MasterTenantID
}
_, err := s.db.ExecContext(ctx, `
INSERT INTO channel_contacts (channel_type, channel_instance, sender_id, user_id, display_name, username, peer_kind, tenant_id)
VALUES ($1, NULLIF($2,''), $3, NULLIF($4,''), NULLIF($5,''), NULLIF($6,''), NULLIF($7,''), $8)
INSERT INTO channel_contacts (channel_type, channel_instance, sender_id, user_id, display_name, username, peer_kind, contact_type, tenant_id)
VALUES ($1, NULLIF($2,''), $3, NULLIF($4,''), NULLIF($5,''), NULLIF($6,''), NULLIF($7,''), $8, $9)
ON CONFLICT (tenant_id, channel_type, sender_id) DO UPDATE SET
display_name = COALESCE(NULLIF($5,''), channel_contacts.display_name),
username = COALESCE(NULLIF($6,''), channel_contacts.username),
user_id = COALESCE(NULLIF($4,''), channel_contacts.user_id),
channel_instance = COALESCE(NULLIF($2,''), channel_contacts.channel_instance),
peer_kind = COALESCE(NULLIF($7,''), channel_contacts.peer_kind),
contact_type = $8,
last_seen_at = NOW()`,
channelType, channelInstance, senderID, userID, displayName, username, peerKind, tenantID,
channelType, channelInstance, senderID, userID, displayName, username, peerKind, contactType, tenantID,
)
return err
}
@@ -88,7 +90,7 @@ func (s *PGContactStore) ListContacts(ctx context.Context, opts store.ContactLis
where, args, argIdx := contactWhereClause(ctx, opts)
query := `SELECT id, channel_type, channel_instance, sender_id, user_id,
display_name, username, avatar_url, peer_kind, merged_id,
display_name, username, avatar_url, peer_kind, contact_type, merged_id,
first_seen_at, last_seen_at
FROM channel_contacts` + where + " ORDER BY last_seen_at DESC"
@@ -116,7 +118,7 @@ func (s *PGContactStore) ListContacts(ctx context.Context, opts store.ContactLis
var c store.ChannelContact
if err := rows.Scan(
&c.ID, &c.ChannelType, &c.ChannelInstance, &c.SenderID, &c.UserID,
&c.DisplayName, &c.Username, &c.AvatarURL, &c.PeerKind, &c.MergedID,
&c.DisplayName, &c.Username, &c.AvatarURL, &c.PeerKind, &c.ContactType, &c.MergedID,
&c.FirstSeenAt, &c.LastSeenAt,
); err != nil {
return nil, err
@@ -147,7 +149,7 @@ func (s *PGContactStore) GetContactsBySenderIDs(ctx context.Context, senderIDs [
query := fmt.Sprintf(`SELECT DISTINCT ON (sender_id)
id, channel_type, channel_instance, sender_id, user_id,
display_name, username, avatar_url, peer_kind, merged_id,
display_name, username, avatar_url, peer_kind, contact_type, merged_id,
first_seen_at, last_seen_at
FROM channel_contacts
WHERE sender_id IN (%s)
@@ -164,7 +166,7 @@ func (s *PGContactStore) GetContactsBySenderIDs(ctx context.Context, senderIDs [
var c store.ChannelContact
if err := rows.Scan(
&c.ID, &c.ChannelType, &c.ChannelInstance, &c.SenderID, &c.UserID,
&c.DisplayName, &c.Username, &c.AvatarURL, &c.PeerKind, &c.MergedID,
&c.DisplayName, &c.Username, &c.AvatarURL, &c.PeerKind, &c.ContactType, &c.MergedID,
&c.FirstSeenAt, &c.LastSeenAt,
); err != nil {
return nil, err
@@ -213,6 +215,9 @@ func (s *PGContactStore) MergeContacts(ctx context.Context, contactIDs []uuid.UU
len(args)-1, inClause, len(args),
)
_, err := s.db.ExecContext(ctx, q, args...)
if err == nil {
s.InvalidateContactResolveCache()
}
return err
}
@@ -236,6 +241,9 @@ func (s *PGContactStore) UnmergeContacts(ctx context.Context, contactIDs []uuid.
inClause, len(args),
)
_, err := s.db.ExecContext(ctx, q, args...)
if err == nil {
s.InvalidateContactResolveCache()
}
return err
}
@@ -243,7 +251,7 @@ func (s *PGContactStore) GetContactsByMergedID(ctx context.Context, mergedID uui
tid := store.TenantIDFromContext(ctx)
q := `SELECT id, channel_type, channel_instance, sender_id, user_id,
display_name, username, avatar_url, peer_kind, merged_id,
display_name, username, avatar_url, peer_kind, contact_type, merged_id,
first_seen_at, last_seen_at
FROM channel_contacts WHERE merged_id = $1 AND tenant_id = $2
ORDER BY last_seen_at DESC`
@@ -259,7 +267,7 @@ func (s *PGContactStore) GetContactsByMergedID(ctx context.Context, mergedID uui
var c store.ChannelContact
if err := rows.Scan(
&c.ID, &c.ChannelType, &c.ChannelInstance, &c.SenderID, &c.UserID,
&c.DisplayName, &c.Username, &c.AvatarURL, &c.PeerKind, &c.MergedID,
&c.DisplayName, &c.Username, &c.AvatarURL, &c.PeerKind, &c.ContactType, &c.MergedID,
&c.FirstSeenAt, &c.LastSeenAt,
); err != nil {
return nil, err
+41 -14
View File
@@ -27,6 +27,7 @@ type permRow struct {
Scope string
ConfigType string
Permission string
UserID string // individual user ID or "*" (group wildcard)
}
// fwCacheEntry holds cached file_writer ConfigPermission rows for a scope.
@@ -76,7 +77,7 @@ func (s *PGConfigPermissionStore) CheckPermission(ctx context.Context, agentID u
s.mu.RLock()
if entry, ok := s.cache[cacheKey]; ok && time.Since(entry.fetched) < permCacheTTL {
s.mu.RUnlock()
return evalPermRows(entry.rows, scope, configType), nil
return evalPermRows(entry.rows, scope, configType, userID), nil
}
s.mu.RUnlock()
@@ -86,8 +87,8 @@ func (s *PGConfigPermissionStore) CheckPermission(ctx context.Context, agentID u
return false, err
}
rows, err := s.db.QueryContext(ctx,
`SELECT scope, config_type, permission FROM agent_config_permissions
WHERE agent_id = $1 AND user_id = $2`+tClause,
`SELECT scope, config_type, permission, user_id FROM agent_config_permissions
WHERE agent_id = $1 AND (user_id = $2 OR user_id = '*')`+tClause,
append([]any{agentID, userID}, tArgs...)...,
)
if err != nil {
@@ -98,7 +99,7 @@ func (s *PGConfigPermissionStore) CheckPermission(ctx context.Context, agentID u
var permRows []permRow
for rows.Next() {
var r permRow
if err := rows.Scan(&r.Scope, &r.ConfigType, &r.Permission); err != nil {
if err := rows.Scan(&r.Scope, &r.ConfigType, &r.Permission, &r.UserID); err != nil {
return false, err
}
permRows = append(permRows, r)
@@ -112,27 +113,53 @@ func (s *PGConfigPermissionStore) CheckPermission(ctx context.Context, agentID u
s.cache[cacheKey] = permCacheEntry{rows: permRows, fetched: time.Now()}
s.mu.Unlock()
return evalPermRows(permRows, scope, configType), nil
return evalPermRows(permRows, scope, configType, userID), nil
}
// evalPermRows evaluates cached permission rows against scope and configType.
func evalPermRows(rows []permRow, scope, configType string) bool {
var hasDeny, hasAllow bool
// Priority-based evaluation: individual permissions override group wildcards (user_id="*").
//
// 1. Individual DENY → REJECT (highest priority)
// 2. Individual ALLOW → ACCEPT
// 3. Group (*) DENY → REJECT
// 4. Group (*) ALLOW → ACCEPT
// 5. No match → REJECT (default)
func evalPermRows(rows []permRow, scope, configType, targetUserID string) bool {
var individualDeny, individualAllow bool
var groupDeny, groupAllow bool
for _, r := range rows {
if !matchWildcard(r.Scope, scope) || !matchWildcard(r.ConfigType, configType) {
continue
}
switch r.Permission {
case "deny":
hasDeny = true
case "allow":
hasAllow = true
if r.UserID == targetUserID {
switch r.Permission {
case "deny":
individualDeny = true
case "allow":
individualAllow = true
}
} else if r.UserID == "*" {
switch r.Permission {
case "deny":
groupDeny = true
case "allow":
groupAllow = true
}
}
}
if hasDeny {
// Individual takes priority over group
if individualDeny {
return false
}
return hasAllow
if individualAllow {
return true
}
if groupDeny {
return false
}
return groupAllow
}
// matchWildcard performs simple wildcard matching for scope/config_type.
+102
View File
@@ -0,0 +1,102 @@
package pg
import (
"context"
"database/sql"
"errors"
"sync"
"time"
"github.com/google/uuid"
"github.com/nextlevelbuilder/goclaw/internal/store"
)
const contactResolveCacheTTL = 60 * time.Second
// contactResolveEntry holds a cached tenant-user resolution result.
type contactResolveEntry struct {
tenantUserID string // empty = not merged
fetched time.Time
}
// contactResolveCache is a TTL cache for contact→tenant-user resolution.
// Mirrors the pattern in config_permissions.go (permCacheTTL).
type contactResolveCache struct {
mu sync.RWMutex
items map[string]contactResolveEntry // key: "tenantID:channelType:senderID"
}
func newContactResolveCache() *contactResolveCache {
return &contactResolveCache{items: make(map[string]contactResolveEntry)}
}
func (c *contactResolveCache) get(key string) (string, bool) {
c.mu.RLock()
defer c.mu.RUnlock()
if entry, ok := c.items[key]; ok && time.Since(entry.fetched) < contactResolveCacheTTL {
return entry.tenantUserID, true
}
return "", false
}
func (c *contactResolveCache) set(key, tenantUserID string) {
c.mu.Lock()
c.items[key] = contactResolveEntry{tenantUserID: tenantUserID, fetched: time.Now()}
c.mu.Unlock()
}
// InvalidateContactResolveCache clears all cached contact→tenant-user resolutions.
// Call after merge/unmerge operations.
func (s *PGContactStore) InvalidateContactResolveCache() {
if s.resolveCache == nil {
return
}
s.resolveCache.mu.Lock()
s.resolveCache.items = make(map[string]contactResolveEntry)
s.resolveCache.mu.Unlock()
}
// ResolveTenantUserID looks up a contact's merged tenant-user identity.
// Uses an in-memory cache with 60s TTL to avoid per-message DB queries.
func (s *PGContactStore) ResolveTenantUserID(ctx context.Context, channelType, senderID string) (string, error) {
tid := store.TenantIDFromContext(ctx)
if tid == uuid.Nil {
return "", nil
}
cacheKey := tid.String() + ":" + channelType + ":" + senderID
// Check cache.
if s.resolveCache != nil {
if resolved, ok := s.resolveCache.get(cacheKey); ok {
return resolved, nil
}
}
// Query DB: join channel_contacts → tenant_users via merged_id.
var tenantUserID string
err := s.db.QueryRowContext(ctx,
`SELECT tu.user_id FROM channel_contacts cc
JOIN tenant_users tu ON cc.merged_id = tu.id
WHERE cc.tenant_id = $1 AND cc.channel_type = $2 AND cc.sender_id = $3
AND cc.merged_id IS NOT NULL`,
tid, channelType, senderID,
).Scan(&tenantUserID)
if errors.Is(err, sql.ErrNoRows) {
// Not merged — cache the negative result too.
if s.resolveCache != nil {
s.resolveCache.set(cacheKey, "")
}
return "", nil
}
if err != nil {
return "", err
}
// Cache positive result.
if s.resolveCache != nil {
s.resolveCache.set(cacheKey, tenantUserID)
}
return tenantUserID, nil
}
+116 -25
View File
@@ -262,46 +262,137 @@ func (s *PGSecureCLIStore) ListByAgent(ctx context.Context, agentID uuid.UUID) (
// LookupByBinary finds the best credential config for a binary name.
// Agent-specific config takes priority over global (agent_id IS NULL).
func (s *PGSecureCLIStore) LookupByBinary(ctx context.Context, binaryName string, agentID *uuid.UUID) (*store.SecureCLIBinary, error) {
var row *sql.Row
// If userID is non-empty, also fetches per-user env overrides via LEFT JOIN (zero extra queries).
func (s *PGSecureCLIStore) LookupByBinary(ctx context.Context, binaryName string, agentID *uuid.UUID, userID string) (*store.SecureCLIBinary, error) {
tid := store.TenantIDFromContext(ctx)
isCross := store.IsCrossTenant(ctx)
if !isCross && tid == uuid.Nil {
return nil, nil
}
// Build query with optional LEFT JOIN for per-user credentials.
selectCols := secureCLISelectCols
joinClause := ""
if userID != "" {
selectCols += ", uc.encrypted_env AS user_env"
joinClause = " LEFT JOIN secure_cli_user_credentials uc ON uc.binary_id = b.id AND uc.user_id = $%d AND uc.tenant_id = $%d"
} else {
selectCols += ", NULL AS user_env"
}
var args []any
var query string
if agentID != nil {
if store.IsCrossTenant(ctx) {
row = s.db.QueryRowContext(ctx,
`SELECT `+secureCLISelectCols+` FROM secure_cli_binaries
WHERE binary_name = $1 AND (agent_id = $2 OR agent_id IS NULL) AND enabled = true
ORDER BY agent_id NULLS LAST LIMIT 1`, binaryName, *agentID)
if userID != "" {
if isCross {
args = []any{binaryName, *agentID, userID, tid}
query = `SELECT ` + selectCols + ` FROM secure_cli_binaries b` +
fmt.Sprintf(joinClause, 3, 4) +
` WHERE b.binary_name = $1 AND (b.agent_id = $2 OR b.agent_id IS NULL) AND b.enabled = true
ORDER BY b.agent_id NULLS LAST LIMIT 1`
} else {
args = []any{binaryName, *agentID, tid, userID, tid}
query = `SELECT ` + selectCols + ` FROM secure_cli_binaries b` +
fmt.Sprintf(joinClause, 4, 5) +
` WHERE b.binary_name = $1 AND (b.agent_id = $2 OR b.agent_id IS NULL) AND b.enabled = true AND b.tenant_id = $3
ORDER BY b.agent_id NULLS LAST LIMIT 1`
}
} else {
tid := store.TenantIDFromContext(ctx)
if tid == uuid.Nil {
return nil, nil
if isCross {
args = []any{binaryName, *agentID}
query = `SELECT ` + selectCols + ` FROM secure_cli_binaries b
WHERE b.binary_name = $1 AND (b.agent_id = $2 OR b.agent_id IS NULL) AND b.enabled = true
ORDER BY b.agent_id NULLS LAST LIMIT 1`
} else {
args = []any{binaryName, *agentID, tid}
query = `SELECT ` + selectCols + ` FROM secure_cli_binaries b
WHERE b.binary_name = $1 AND (b.agent_id = $2 OR b.agent_id IS NULL) AND b.enabled = true AND b.tenant_id = $3
ORDER BY b.agent_id NULLS LAST LIMIT 1`
}
row = s.db.QueryRowContext(ctx,
`SELECT `+secureCLISelectCols+` FROM secure_cli_binaries
WHERE binary_name = $1 AND (agent_id = $2 OR agent_id IS NULL) AND enabled = true AND tenant_id = $3
ORDER BY agent_id NULLS LAST LIMIT 1`, binaryName, *agentID, tid)
}
} else {
if store.IsCrossTenant(ctx) {
row = s.db.QueryRowContext(ctx,
`SELECT `+secureCLISelectCols+` FROM secure_cli_binaries
WHERE binary_name = $1 AND agent_id IS NULL AND enabled = true LIMIT 1`, binaryName)
if userID != "" {
if isCross {
args = []any{binaryName, userID, tid}
query = `SELECT ` + selectCols + ` FROM secure_cli_binaries b` +
fmt.Sprintf(joinClause, 2, 3) +
` WHERE b.binary_name = $1 AND b.agent_id IS NULL AND b.enabled = true LIMIT 1`
} else {
args = []any{binaryName, tid, userID, tid}
query = `SELECT ` + selectCols + ` FROM secure_cli_binaries b` +
fmt.Sprintf(joinClause, 3, 4) +
` WHERE b.binary_name = $1 AND b.agent_id IS NULL AND b.enabled = true AND b.tenant_id = $2 LIMIT 1`
}
} else {
tid := store.TenantIDFromContext(ctx)
if tid == uuid.Nil {
return nil, nil
if isCross {
args = []any{binaryName}
query = `SELECT ` + selectCols + ` FROM secure_cli_binaries b
WHERE b.binary_name = $1 AND b.agent_id IS NULL AND b.enabled = true LIMIT 1`
} else {
args = []any{binaryName, tid}
query = `SELECT ` + selectCols + ` FROM secure_cli_binaries b
WHERE b.binary_name = $1 AND b.agent_id IS NULL AND b.enabled = true AND b.tenant_id = $2 LIMIT 1`
}
row = s.db.QueryRowContext(ctx,
`SELECT `+secureCLISelectCols+` FROM secure_cli_binaries
WHERE binary_name = $1 AND agent_id IS NULL AND enabled = true AND tenant_id = $2 LIMIT 1`, binaryName, tid)
}
}
b, err := s.scanRow(row)
row := s.db.QueryRowContext(ctx, query, args...)
b, err := s.scanRowWithUserEnv(row)
if err == sql.ErrNoRows {
return nil, nil
}
return b, err
}
// scanRowWithUserEnv scans a row that includes the extra user_env column from LEFT JOIN.
func (s *PGSecureCLIStore) scanRowWithUserEnv(row *sql.Row) (*store.SecureCLIBinary, error) {
var b store.SecureCLIBinary
var binaryPath *string
var agentID *uuid.UUID
var denyArgs, denyVerbose *[]byte
var env []byte
var userEnv []byte
err := row.Scan(
&b.ID, &b.BinaryName, &binaryPath, &b.Description, &env,
&denyArgs, &denyVerbose,
&b.TimeoutSeconds, &b.Tips, &agentID,
&b.Enabled, &b.CreatedBy, &b.CreatedAt, &b.UpdatedAt,
&userEnv,
)
if err != nil {
return nil, err
}
b.BinaryPath = binaryPath
b.AgentID = agentID
if denyArgs != nil {
b.DenyArgs = *denyArgs
}
if denyVerbose != nil {
b.DenyVerbose = *denyVerbose
}
// Decrypt base env
if len(env) > 0 && s.encKey != "" {
if decrypted, err := crypto.Decrypt(string(env), s.encKey); err == nil {
b.EncryptedEnv = []byte(decrypted)
}
} else {
b.EncryptedEnv = env
}
// Decrypt per-user env
if len(userEnv) > 0 && s.encKey != "" {
if decrypted, err := crypto.Decrypt(string(userEnv), s.encKey); err == nil {
b.UserEnv = []byte(decrypted)
}
}
return &b, nil
}
func (s *PGSecureCLIStore) ListEnabled(ctx context.Context) ([]store.SecureCLIBinary, error) {
query := `SELECT ` + secureCLISelectCols + ` FROM secure_cli_binaries WHERE enabled = true`
var qArgs []any
@@ -0,0 +1,119 @@
package pg
import (
"context"
"database/sql"
"errors"
"fmt"
"time"
"github.com/google/uuid"
"github.com/nextlevelbuilder/goclaw/internal/crypto"
"github.com/nextlevelbuilder/goclaw/internal/store"
)
func (s *PGSecureCLIStore) GetUserCredentials(ctx context.Context, binaryID uuid.UUID, userID string) (*store.SecureCLIUserCredential, error) {
tid := store.TenantIDFromContext(ctx)
if tid == uuid.Nil {
tid = store.MasterTenantID
}
var uc store.SecureCLIUserCredential
var env []byte
err := s.db.QueryRowContext(ctx,
`SELECT id, binary_id, user_id, encrypted_env, metadata, created_at, updated_at
FROM secure_cli_user_credentials
WHERE binary_id = $1 AND user_id = $2 AND tenant_id = $3`,
binaryID, userID, tid,
).Scan(&uc.ID, &uc.BinaryID, &uc.UserID, &env, &uc.Metadata, &uc.CreatedAt, &uc.UpdatedAt)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
if err != nil {
return nil, err
}
// Decrypt env
if len(env) > 0 && s.encKey != "" {
if decrypted, err := crypto.Decrypt(string(env), s.encKey); err == nil {
uc.EncryptedEnv = []byte(decrypted)
}
} else {
uc.EncryptedEnv = env
}
return &uc, nil
}
func (s *PGSecureCLIStore) SetUserCredentials(ctx context.Context, binaryID uuid.UUID, userID string, encryptedEnv []byte) error {
tid := store.TenantIDFromContext(ctx)
if tid == uuid.Nil {
tid = store.MasterTenantID
}
// Encrypt env
var envBytes []byte
if len(encryptedEnv) > 0 && s.encKey != "" {
encrypted, err := crypto.Encrypt(string(encryptedEnv), s.encKey)
if err != nil {
return fmt.Errorf("encrypt env: %w", err)
}
envBytes = []byte(encrypted)
} else {
envBytes = encryptedEnv
}
now := time.Now()
_, err := s.db.ExecContext(ctx,
`INSERT INTO secure_cli_user_credentials (binary_id, user_id, encrypted_env, metadata, tenant_id, created_at, updated_at)
VALUES ($1, $2, $3, '{}', $4, $5, $5)
ON CONFLICT (binary_id, user_id, tenant_id) DO UPDATE SET
encrypted_env = EXCLUDED.encrypted_env,
updated_at = EXCLUDED.updated_at`,
binaryID, userID, envBytes, tid, now,
)
return err
}
func (s *PGSecureCLIStore) DeleteUserCredentials(ctx context.Context, binaryID uuid.UUID, userID string) error {
tid := store.TenantIDFromContext(ctx)
if tid == uuid.Nil {
tid = store.MasterTenantID
}
_, err := s.db.ExecContext(ctx,
`DELETE FROM secure_cli_user_credentials WHERE binary_id = $1 AND user_id = $2 AND tenant_id = $3`,
binaryID, userID, tid,
)
return err
}
func (s *PGSecureCLIStore) ListUserCredentials(ctx context.Context, binaryID uuid.UUID) ([]store.SecureCLIUserCredential, error) {
tid := store.TenantIDFromContext(ctx)
if tid == uuid.Nil {
tid = store.MasterTenantID
}
rows, err := s.db.QueryContext(ctx,
`SELECT id, binary_id, user_id, encrypted_env, metadata, created_at, updated_at
FROM secure_cli_user_credentials
WHERE binary_id = $1 AND tenant_id = $2
ORDER BY created_at`, binaryID, tid)
if err != nil {
return nil, err
}
defer rows.Close()
var result []store.SecureCLIUserCredential
for rows.Next() {
var uc store.SecureCLIUserCredential
var env []byte
if err := rows.Scan(&uc.ID, &uc.BinaryID, &uc.UserID, &env, &uc.Metadata, &uc.CreatedAt, &uc.UpdatedAt); err != nil {
return nil, err
}
if len(env) > 0 && s.encKey != "" {
if decrypted, err := crypto.Decrypt(string(env), s.encKey); err == nil {
uc.EncryptedEnv = []byte(decrypted)
}
} else {
uc.EncryptedEnv = env
}
result = append(result, uc)
}
return result, rows.Err()
}
+23 -1
View File
@@ -22,6 +22,19 @@ type SecureCLIBinary struct {
AgentID *uuid.UUID `json:"agent_id,omitempty"`
Enabled bool `json:"enabled"`
CreatedBy string `json:"created_by"`
UserEnv []byte `json:"-"` // per-user encrypted env (populated by LookupByBinary LEFT JOIN)
}
// SecureCLIUserCredential holds per-user encrypted env overrides for a binary.
type SecureCLIUserCredential struct {
ID uuid.UUID `json:"id"`
BinaryID uuid.UUID `json:"binary_id"`
UserID string `json:"user_id"`
Metadata json.RawMessage `json:"metadata,omitempty"`
CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"`
// EncryptedEnv is decrypted JSON — never serialized to API.
EncryptedEnv []byte `json:"-"`
}
// SecureCLIStore manages secure CLI binary credential configurations.
@@ -35,8 +48,17 @@ type SecureCLIStore interface {
// LookupByBinary finds the best-matching credential config for a binary name.
// Priority: agent-specific > global (agent_id IS NULL). Returns nil if not found.
LookupByBinary(ctx context.Context, binaryName string, agentID *uuid.UUID) (*SecureCLIBinary, error)
// If userID is non-empty, also fetches per-user env overrides via LEFT JOIN
// and populates SecureCLIBinary.UserEnv (zero extra queries).
LookupByBinary(ctx context.Context, binaryName string, agentID *uuid.UUID, userID string) (*SecureCLIBinary, error)
// ListEnabled returns all enabled configs (for TOOLS.md context generation).
ListEnabled(ctx context.Context) ([]SecureCLIBinary, error)
// --- Per-user credential management ---
GetUserCredentials(ctx context.Context, binaryID uuid.UUID, userID string) (*SecureCLIUserCredential, error)
SetUserCredentials(ctx context.Context, binaryID uuid.UUID, userID string, encryptedEnv []byte) error
DeleteUserCredentials(ctx context.Context, binaryID uuid.UUID, userID string) error
ListUserCredentials(ctx context.Context, binaryID uuid.UUID) ([]SecureCLIUserCredential, error)
}
@@ -21,6 +21,7 @@ type permRow struct {
Scope string
ConfigType string
Permission string
UserID string
}
type permCacheEntry struct {
@@ -72,7 +73,7 @@ func (s *SQLiteConfigPermissionStore) CheckPermission(ctx context.Context, agent
s.mu.RLock()
if entry, ok := s.cache[cacheKey]; ok && time.Since(entry.fetched) < permCacheTTL {
s.mu.RUnlock()
return evalPermRows(entry.rows, scope, configType), nil
return evalPermRows(entry.rows, scope, configType, userID), nil
}
s.mu.RUnlock()
@@ -81,8 +82,8 @@ func (s *SQLiteConfigPermissionStore) CheckPermission(ctx context.Context, agent
return false, err
}
rows, err := s.db.QueryContext(ctx,
`SELECT scope, config_type, permission FROM agent_config_permissions
WHERE agent_id = ? AND user_id = ?`+tClause,
`SELECT scope, config_type, permission, user_id FROM agent_config_permissions
WHERE agent_id = ? AND (user_id = ? OR user_id = '*')`+tClause,
append([]any{agentID, userID}, tArgs...)...,
)
if err != nil {
@@ -93,7 +94,7 @@ func (s *SQLiteConfigPermissionStore) CheckPermission(ctx context.Context, agent
var permRows []permRow
for rows.Next() {
var r permRow
if err := rows.Scan(&r.Scope, &r.ConfigType, &r.Permission); err != nil {
if err := rows.Scan(&r.Scope, &r.ConfigType, &r.Permission, &r.UserID); err != nil {
return false, err
}
permRows = append(permRows, r)
@@ -106,7 +107,7 @@ func (s *SQLiteConfigPermissionStore) CheckPermission(ctx context.Context, agent
s.cache[cacheKey] = permCacheEntry{rows: permRows, fetched: time.Now()}
s.mu.Unlock()
return evalPermRows(permRows, scope, configType), nil
return evalPermRows(permRows, scope, configType, userID), nil
}
func (s *SQLiteConfigPermissionStore) Grant(ctx context.Context, perm *store.ConfigPermission) error {
@@ -243,24 +244,43 @@ func scanConfigPermissions(rows *sql.Rows) ([]store.ConfigPermission, error) {
return perms, nil
}
// evalPermRows evaluates cached permission rows against scope and configType (deny-first).
func evalPermRows(rows []permRow, scope, configType string) bool {
var hasDeny, hasAllow bool
// evalPermRows evaluates cached permission rows with priority-based evaluation.
// Individual permissions (matching targetUserID) override group wildcards (user_id="*").
func evalPermRows(rows []permRow, scope, configType, targetUserID string) bool {
var individualDeny, individualAllow bool
var groupDeny, groupAllow bool
for _, r := range rows {
if !matchWildcard(r.Scope, scope) || !matchWildcard(r.ConfigType, configType) {
continue
}
switch r.Permission {
case "deny":
hasDeny = true
case "allow":
hasAllow = true
if r.UserID == targetUserID {
switch r.Permission {
case "deny":
individualDeny = true
case "allow":
individualAllow = true
}
} else if r.UserID == "*" {
switch r.Permission {
case "deny":
groupDeny = true
case "allow":
groupAllow = true
}
}
}
if hasDeny {
if individualDeny {
return false
}
return hasAllow
if individualAllow {
return true
}
if groupDeny {
return false
}
return groupAllow
}
// matchWildcard performs simple wildcard matching for scope/config_type.
+26 -9
View File
@@ -5,6 +5,7 @@ package sqlitestore
import (
"context"
"database/sql"
"errors"
"fmt"
"strings"
@@ -22,22 +23,21 @@ func NewSQLiteContactStore(db *sql.DB) *SQLiteContactStore {
return &SQLiteContactStore{db: db}
}
func (s *SQLiteContactStore) UpsertContact(ctx context.Context, channelType, channelInstance, senderID, userID, displayName, username, peerKind string) error {
func (s *SQLiteContactStore) UpsertContact(ctx context.Context, channelType, channelInstance, senderID, userID, displayName, username, peerKind, contactType string) error {
tenantID := store.TenantIDFromContext(ctx)
if tenantID == uuid.Nil {
tenantID = store.MasterTenantID
}
// SQLite UPSERT: ON CONFLICT ... DO UPDATE SET
// NULLIF emulated with CASE WHEN
_, err := s.db.ExecContext(ctx, `
INSERT INTO channel_contacts (channel_type, channel_instance, sender_id, user_id, display_name, username, peer_kind, tenant_id)
VALUES (?, NULLIF(?,?), ?, NULLIF(?,?), NULLIF(?,?), NULLIF(?,?), NULLIF(?,?), ?)
INSERT INTO channel_contacts (channel_type, channel_instance, sender_id, user_id, display_name, username, peer_kind, contact_type, tenant_id)
VALUES (?, NULLIF(?,?), ?, NULLIF(?,?), NULLIF(?,?), NULLIF(?,?), NULLIF(?,?), ?, ?)
ON CONFLICT (tenant_id, channel_type, sender_id) DO UPDATE SET
display_name = COALESCE(NULLIF(excluded.display_name,''), channel_contacts.display_name),
username = COALESCE(NULLIF(excluded.username,''), channel_contacts.username),
user_id = COALESCE(NULLIF(excluded.user_id,''), channel_contacts.user_id),
channel_instance = COALESCE(NULLIF(excluded.channel_instance,''), channel_contacts.channel_instance),
peer_kind = COALESCE(NULLIF(excluded.peer_kind,''), channel_contacts.peer_kind),
contact_type = excluded.contact_type,
last_seen_at = CURRENT_TIMESTAMP`,
channelType,
channelInstance, "",
@@ -46,6 +46,7 @@ func (s *SQLiteContactStore) UpsertContact(ctx context.Context, channelType, cha
displayName, "",
username, "",
peerKind, "",
contactType,
tenantID,
)
return err
@@ -86,14 +87,14 @@ func contactWhereSQLite(ctx context.Context, opts store.ContactListOpts) (string
}
const contactSelectCols = `id, channel_type, channel_instance, sender_id, user_id,
display_name, username, avatar_url, peer_kind, merged_id,
display_name, username, avatar_url, peer_kind, contact_type, merged_id,
first_seen_at, last_seen_at`
func scanContact(rows *sql.Rows) (store.ChannelContact, error) {
var c store.ChannelContact
err := rows.Scan(
&c.ID, &c.ChannelType, &c.ChannelInstance, &c.SenderID, &c.UserID,
&c.DisplayName, &c.Username, &c.AvatarURL, &c.PeerKind, &c.MergedID,
&c.DisplayName, &c.Username, &c.AvatarURL, &c.PeerKind, &c.ContactType, &c.MergedID,
&c.FirstSeenAt, &c.LastSeenAt,
)
return c, err
@@ -150,7 +151,7 @@ func (s *SQLiteContactStore) GetContactsBySenderIDs(ctx context.Context, senderI
// SQLite has no DISTINCT ON; emulate with GROUP BY + MAX rowid trick via subquery
query := `SELECT id, channel_type, channel_instance, sender_id, user_id,
display_name, username, avatar_url, peer_kind, merged_id,
display_name, username, avatar_url, peer_kind, contact_type, merged_id,
first_seen_at, last_seen_at
FROM channel_contacts
WHERE sender_id IN (` + strings.Join(placeholders, ",") + `)
@@ -182,7 +183,7 @@ func (s *SQLiteContactStore) GetContactByID(ctx context.Context, id uuid.UUID) (
var c store.ChannelContact
if err := row.Scan(
&c.ID, &c.ChannelType, &c.ChannelInstance, &c.SenderID, &c.UserID,
&c.DisplayName, &c.Username, &c.AvatarURL, &c.PeerKind, &c.MergedID,
&c.DisplayName, &c.Username, &c.AvatarURL, &c.PeerKind, &c.ContactType, &c.MergedID,
&c.FirstSeenAt, &c.LastSeenAt,
); err != nil {
return nil, err
@@ -253,3 +254,19 @@ func (s *SQLiteContactStore) GetContactsByMergedID(ctx context.Context, mergedID
}
return contacts, rows.Err()
}
func (s *SQLiteContactStore) ResolveTenantUserID(ctx context.Context, channelType, senderID string) (string, error) {
tid := store.TenantIDFromContext(ctx)
var tenantUserID string
err := s.db.QueryRowContext(ctx,
`SELECT tu.user_id FROM channel_contacts cc
JOIN tenant_users tu ON cc.merged_id = tu.id
WHERE cc.tenant_id = ? AND cc.channel_type = ? AND cc.sender_id = ?
AND cc.merged_id IS NOT NULL`,
tid, channelType, senderID,
).Scan(&tenantUserID)
if errors.Is(err, sql.ErrNoRows) {
return "", nil
}
return tenantUserID, err
}
+3 -2
View File
@@ -14,7 +14,7 @@ var schemaSQL string
// SchemaVersion is the current SQLite schema version.
// Bump this when adding new migration steps below.
const SchemaVersion = 1
const SchemaVersion = 2
// migrations maps version → SQL to apply when upgrading FROM that version.
// schema.sql always represents the LATEST full schema (for fresh DBs).
@@ -28,7 +28,8 @@ const SchemaVersion = 1
//
// Then bump SchemaVersion to 2.
var migrations = map[int]string{
// Version 1 is the initial schema — no patch needed (schema.sql covers it).
// Version 1 → 2: add contact_type column to channel_contacts.
1: `ALTER TABLE channel_contacts ADD COLUMN contact_type VARCHAR(20) NOT NULL DEFAULT 'user';`,
}
// EnsureSchema creates tables if they don't exist and applies incremental migrations.
+1
View File
@@ -1036,6 +1036,7 @@ CREATE TABLE IF NOT EXISTS channel_contacts (
username VARCHAR(255),
avatar_url TEXT,
peer_kind VARCHAR(20),
contact_type VARCHAR(20) NOT NULL DEFAULT 'user',
metadata TEXT DEFAULT '{}',
merged_id TEXT,
tenant_id TEXT NOT NULL REFERENCES tenants(id),
+12 -1
View File
@@ -6,6 +6,7 @@ import (
"encoding/json"
"fmt"
"log/slog"
"maps"
"os"
"os/exec"
"regexp"
@@ -136,6 +137,14 @@ func (t *ExecTool) executeCredentialed(ctx context.Context, cred *store.SecureCL
}
}
// Step 4b: Merge per-user env overrides (user takes priority over base)
if len(cred.UserEnv) > 0 {
var userEnvMap map[string]string
if err := json.Unmarshal(cred.UserEnv, &userEnvMap); err == nil {
maps.Copy(envMap, userEnvMap)
}
}
// Step 5: Register credential values for output scrubbing
for _, v := range envMap {
AddCredentialScrubValues(v)
@@ -275,7 +284,9 @@ func (t *ExecTool) lookupCredentialedBinary(ctx context.Context, command string)
if agentID != uuid.Nil {
agentIDPtr = &agentID
}
cred, err := t.secureCLIStore.LookupByBinary(ctx, binary, agentIDPtr)
// Pass userID for per-user credential resolution (LEFT JOIN, zero extra queries).
userID := store.UserIDFromContext(ctx)
cred, err := t.secureCLIStore.LookupByBinary(ctx, binary, agentIDPtr, userID)
if err != nil || cred == nil {
return nil, "", nil
}
+9 -11
View File
@@ -6,6 +6,7 @@ import (
"encoding/json"
"errors"
"fmt"
"strings"
"time"
"github.com/google/uuid"
@@ -268,10 +269,13 @@ func (t *HeartbeatTool) checkPermission(ctx context.Context, agentID uuid.UUID)
if t.permStore == nil {
return nil // no permission store = allow all (backward compat)
}
userID := store.UserIDFromContext(ctx)
if userID == "" {
// Use individual sender ID for permission check (not group-scoped UserID).
// This matches the pattern used by file_writer and /addwriter commands.
senderID := store.SenderIDFromContext(ctx)
if senderID == "" {
return nil // system context (cron, subagent) = allow
}
numericID := strings.SplitN(senderID, "|", 2)[0]
// Determine scope from context: "agent" for DM, "group:{channel}:{chatId}" for groups.
scope := "agent"
@@ -281,7 +285,7 @@ func (t *HeartbeatTool) checkPermission(ctx context.Context, agentID uuid.UUID)
}
}
allowed, err := t.permStore.CheckPermission(ctx, agentID, scope, "heartbeat", userID)
allowed, err := t.permStore.CheckPermission(ctx, agentID, scope, "heartbeat", numericID)
if err != nil {
return fmt.Errorf("permission check failed: %w", err)
}
@@ -292,14 +296,8 @@ func (t *HeartbeatTool) checkPermission(ctx context.Context, agentID uuid.UUID)
// Fallback: check if user is agent owner (via agent store).
if t.agentStore != nil {
ag, agErr := t.agentStore.GetByID(ctx, agentID)
if agErr == nil {
senderID := store.SenderIDFromContext(ctx)
if senderID == "" {
senderID = userID
}
if ag.OwnerID != "" && ag.OwnerID == senderID {
return nil // agent owner = allow
}
if agErr == nil && ag.OwnerID != "" && ag.OwnerID == senderID {
return nil // agent owner = allow
}
}
+7 -7
View File
@@ -13,8 +13,8 @@ import (
// ── mock KG store ──────────────────────────────────────────────────
type mockKGStore struct {
entities map[string]store.Entity // keyed by entity ID
relations []store.Relation // all relations
entities map[string]store.Entity // keyed by entity ID
relations []store.Relation // all relations
traversal map[string][]store.TraversalResult // keyed by start entity ID
}
@@ -122,7 +122,7 @@ func (m *mockKGStore) Stats(context.Context, string, string) (*store.GraphStats,
}
func (m *mockKGStore) SetEmbeddingProvider(store.EmbeddingProvider) {}
func (m *mockKGStore) Close() error { return nil }
func (m *mockKGStore) Close() error { return nil }
// ── test helpers ───────────────────────────────────────────────────
@@ -174,7 +174,7 @@ func setupBaseGraph() (*mockKGStore, map[string]string) {
// Pre-compute traversal results (mock outgoing-only behavior)
// A → B, C (outgoing from A)
ms.traversal[ids["A"]] = []store.TraversalResult{
{Entity: entities[1], Depth: 1, Via: "owns"}, // B=GoClaw
{Entity: entities[1], Depth: 1, Via: "owns"}, // B=GoClaw
{Entity: entities[2], Depth: 1, Via: "manages"}, // C=Dầu thô
}
// B → C (outgoing from B)
@@ -273,7 +273,7 @@ func TestKGTraversal_Tier2_CappedAt10(t *testing.T) {
ms.entities[entityX] = store.Entity{ID: entityX, AgentID: testAgentID.String(), UserID: testUserID, Name: "HubNode", EntityType: "concept"}
// Create 15 incoming relations to X
for i := 0; i < 15; i++ {
for i := range 15 {
srcID := uuid.NewString()
srcName := fmt.Sprintf("Source_%02d", i)
ms.entities[srcID] = store.Entity{ID: srcID, AgentID: testAgentID.String(), UserID: testUserID, Name: srcName, EntityType: "concept"}
@@ -359,7 +359,7 @@ func TestKGTraversal_Tier1_CappedAt20(t *testing.T) {
// Create 25 traversal results
var results []store.TraversalResult
for i := 0; i < 25; i++ {
for i := range 25 {
eid := uuid.NewString()
name := fmt.Sprintf("Node_%02d", i)
ms.entities[eid] = store.Entity{ID: eid, AgentID: testAgentID.String(), UserID: testUserID, Name: name, EntityType: "concept"}
@@ -389,7 +389,7 @@ func TestKGSearch_RelationsCappedAt5(t *testing.T) {
ms.entities[entityID] = store.Entity{ID: entityID, AgentID: testAgentID.String(), UserID: testUserID, Name: "HubEntity", EntityType: "concept"}
// Create 8 outgoing relations from entity
for i := 0; i < 8; i++ {
for i := range 8 {
tgtID := uuid.NewString()
tgtName := fmt.Sprintf("Target_%02d", i)
ms.entities[tgtID] = store.Entity{ID: tgtID, AgentID: testAgentID.String(), UserID: testUserID, Name: tgtName, EntityType: "concept"}
+3 -3
View File
@@ -46,8 +46,8 @@ func (m *mockMemoryStore) ListDocuments(_ context.Context, agentID, userID strin
var out []store.DocumentInfo
prefix := agentID + "|" + userID + "|"
for k := range m.docs {
if strings.HasPrefix(k, prefix) {
path := strings.TrimPrefix(k, prefix)
if after, ok := strings.CutPrefix(k, prefix); ok {
path := after
out = append(out, store.DocumentInfo{Path: path})
}
}
@@ -72,7 +72,7 @@ func (m *mockMemoryStore) Search(_ context.Context, _ string, _, _ string, _ sto
}
func (m *mockMemoryStore) IndexDocument(_ context.Context, _, _, _ string) error { return nil }
func (m *mockMemoryStore) IndexAll(_ context.Context, _, _ string) error { return nil }
func (m *mockMemoryStore) SetEmbeddingProvider(_ store.EmbeddingProvider) {}
func (m *mockMemoryStore) SetEmbeddingProvider(_ store.EmbeddingProvider) {}
func (m *mockMemoryStore) Close() error { return nil }
// --- Test helpers ---
+18 -18
View File
@@ -111,25 +111,25 @@ func (m *mockSessionStore) SetLabel(_ context.Context, key, label string) {
}
func (m *mockSessionStore) SetAgentInfo(context.Context, string, uuid.UUID, string) {}
func (m *mockSessionStore) TruncateHistory(context.Context, string, int) {}
func (m *mockSessionStore) SetHistory(context.Context, string, []providers.Message) {}
func (m *mockSessionStore) Reset(context.Context, string) {}
func (m *mockSessionStore) Delete(context.Context, string) error { return nil }
func (m *mockSessionStore) Save(context.Context, string) error { return nil }
func (m *mockSessionStore) TruncateHistory(context.Context, string, int) {}
func (m *mockSessionStore) SetHistory(context.Context, string, []providers.Message) {}
func (m *mockSessionStore) Reset(context.Context, string) {}
func (m *mockSessionStore) Delete(context.Context, string) error { return nil }
func (m *mockSessionStore) Save(context.Context, string) error { return nil }
func (m *mockSessionStore) UpdateMetadata(context.Context, string, string, string, string) {}
func (m *mockSessionStore) AccumulateTokens(context.Context, string, int64, int64) {}
func (m *mockSessionStore) IncrementCompaction(context.Context, string) {}
func (m *mockSessionStore) GetCompactionCount(context.Context, string) int { return 0 }
func (m *mockSessionStore) GetMemoryFlushCompactionCount(context.Context, string) int { return 0 }
func (m *mockSessionStore) SetMemoryFlushDone(context.Context, string) {}
func (m *mockSessionStore) GetSessionMetadata(context.Context, string) map[string]string { return nil }
func (m *mockSessionStore) SetSessionMetadata(context.Context, string, map[string]string) {}
func (m *mockSessionStore) SetSpawnInfo(context.Context, string, string, int) {}
func (m *mockSessionStore) SetContextWindow(context.Context, string, int) {}
func (m *mockSessionStore) GetContextWindow(context.Context, string) int { return 0 }
func (m *mockSessionStore) SetLastPromptTokens(context.Context, string, int, int) {}
func (m *mockSessionStore) GetLastPromptTokens(context.Context, string) (int, int) { return 0, 0 }
func (m *mockSessionStore) IncrementCompaction(context.Context, string) {}
func (m *mockSessionStore) GetCompactionCount(context.Context, string) int { return 0 }
func (m *mockSessionStore) GetMemoryFlushCompactionCount(context.Context, string) int { return 0 }
func (m *mockSessionStore) SetMemoryFlushDone(context.Context, string) {}
func (m *mockSessionStore) GetSessionMetadata(context.Context, string) map[string]string { return nil }
func (m *mockSessionStore) SetSessionMetadata(context.Context, string, map[string]string) {}
func (m *mockSessionStore) SetSpawnInfo(context.Context, string, string, int) {}
func (m *mockSessionStore) SetContextWindow(context.Context, string, int) {}
func (m *mockSessionStore) GetContextWindow(context.Context, string) int { return 0 }
func (m *mockSessionStore) SetLastPromptTokens(context.Context, string, int, int) {}
func (m *mockSessionStore) GetLastPromptTokens(context.Context, string) (int, int) { return 0, 0 }
func (m *mockSessionStore) List(_ context.Context, agentID string) []store.SessionInfo {
m.mu.RLock()
@@ -230,7 +230,7 @@ func TestSessionsList_ActiveMinutesFilter(t *testing.T) {
func TestSessionsList_LimitCapsResults(t *testing.T) {
ms := newMockSessionStore()
for i := 0; i < 5; i++ {
for i := range 5 {
ms.seed("agent:"+sessTestAgentID+":ws:direct:"+string(rune('a'+i)), nil, "")
}
@@ -408,7 +408,7 @@ func TestSessionsHistory_LimitFromEnd(t *testing.T) {
ms := newMockSessionStore()
key := "agent:" + sessTestAgentID + ":ws:direct:1"
var msgs []providers.Message
for i := 0; i < 10; i++ {
for range 10 {
msgs = append(msgs, providers.Message{Role: "user", Content: "msg"})
}
ms.seed(key, msgs, "")
+3 -6
View File
@@ -100,7 +100,7 @@ func estimateArraySize(arr []json.RawMessage, elemLens []int, keep int) int {
placeholderLen := 22 + digitCount(dropped) + 17 // key + number + " elements omitted"}
total := 2 // [ and ]
for i := 0; i < keepHead; i++ {
for i := range keepHead {
total += elemLens[i]
}
total += placeholderLen
@@ -131,7 +131,7 @@ func buildTruncatedArray(arr []json.RawMessage, keepHead, keepTail int) string {
placeholder := fmt.Sprintf(`{"__truncated__":"%d elements omitted"}`, dropped)
parts := make([]string, 0, keepHead+1+keepTail)
for i := 0; i < keepHead; i++ {
for i := range keepHead {
parts = append(parts, string(arr[i]))
}
parts = append(parts, placeholder)
@@ -147,10 +147,7 @@ func truncateArrayElements(arr []json.RawMessage, maxLen int) string {
dropped := len(arr) - 2
placeholder := fmt.Sprintf(`{"__truncated__":"%d elements omitted"}`, dropped)
overhead := len("[,,]") + len(placeholder)
perElem := (maxLen - overhead) / 2
if perElem < 50 {
perElem = 50
}
perElem := max((maxLen-overhead)/2, 50)
first := TruncateMid(string(arr[0]), perElem)
last := TruncateMid(string(arr[len(arr)-1]), perElem)
+1 -1
View File
@@ -341,7 +341,7 @@ func removeQuarantine(appPath string) {
func isNewer(a, b string) bool {
pa := parseSemver(a)
pb := parseSemver(b)
for i := 0; i < 3; i++ {
for i := range 3 {
if pa[i] > pb[i] {
return true
}
+1 -1
View File
@@ -2,4 +2,4 @@ package upgrade
// RequiredSchemaVersion is the schema migration version this binary requires.
// Bump this whenever adding a new SQL migration file.
const RequiredSchemaVersion uint = 31
const RequiredSchemaVersion uint = 32
@@ -0,0 +1,2 @@
ALTER TABLE channel_contacts DROP COLUMN IF EXISTS contact_type;
DROP TABLE IF EXISTS secure_cli_user_credentials;
@@ -0,0 +1,20 @@
-- Per-user credentials for secure CLI binaries.
-- Mirrors mcp_user_credentials pattern: user-specific env vars override binary defaults.
CREATE TABLE secure_cli_user_credentials (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
binary_id UUID NOT NULL REFERENCES secure_cli_binaries(id) ON DELETE CASCADE,
user_id VARCHAR(255) NOT NULL,
encrypted_env BYTEA NOT NULL, -- AES-256-GCM encrypted JSON: {"GH_TOKEN":"xxx"}
metadata JSONB NOT NULL DEFAULT '{}',
tenant_id UUID NOT NULL REFERENCES tenants(id),
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
UNIQUE(binary_id, user_id, tenant_id)
);
CREATE INDEX idx_scuc_tenant ON secure_cli_user_credentials(tenant_id);
CREATE INDEX idx_scuc_binary ON secure_cli_user_credentials(binary_id);
-- Add contact_type column to channel_contacts to distinguish user vs group contacts.
-- Default "user" for backward compatibility with existing records.
ALTER TABLE channel_contacts ADD COLUMN IF NOT EXISTS contact_type VARCHAR(20) NOT NULL DEFAULT 'user';
@@ -0,0 +1,54 @@
import { Combobox } from "@/components/ui/combobox";
import { useUserPicker } from "@/hooks/use-user-picker";
interface UserPickerComboboxProps {
value: string;
onChange: (value: string) => void;
placeholder?: string;
className?: string;
/** Filter contacts by peer_kind: "direct" | "group" | undefined (all). */
peerKind?: "direct" | "group";
/** Filter by source: "contact" | "tenant_user" | undefined (both).
* Use "tenant_user" for merge dialogs and tenant user pickers. */
source?: "contact" | "tenant_user";
/** Allow typing custom values not in the list. Default true. */
allowCustom?: boolean;
/** Render dropdown into a portal container (useful inside dialogs). */
portalContainer?: React.RefObject<HTMLElement | null>;
}
/**
* Unified user picker that searches both channel_contacts and tenant_users.
* Drop-in replacement for Combobox + useContactSearch/useContactPicker.
*
* - Shows 30 most recent results when opened (no typing needed)
* - Debounced server-side search as user types
* - Source badges: [telegram], [discord], [tenant], merged status
* - Deduplicates merged contacts
*
* Uses `value` prop as search term (same pattern as useContactSearch(userId)).
*/
export function UserPickerCombobox({
value,
onChange,
placeholder,
className,
peerKind,
source,
allowCustom = true,
portalContainer,
}: UserPickerComboboxProps) {
const { options } = useUserPicker(value, peerKind, source);
return (
<Combobox
value={value}
onChange={onChange}
options={options}
placeholder={placeholder}
className={className}
allowCustom={allowCustom}
portalContainer={portalContainer}
/>
);
}
+70
View File
@@ -0,0 +1,70 @@
import { useState, useEffect, useMemo } from "react";
import { useQuery } from "@tanstack/react-query";
import { useHttp } from "@/hooks/use-ws";
import { queryKeys } from "@/lib/query-keys";
import type { ComboboxOption } from "@/components/ui/combobox";
/** Unified search result from contacts + tenant_users. */
export interface UserPickerItem {
id: string;
display_name?: string;
username?: string;
source: "contact" | "tenant_user";
channel_type?: string;
peer_kind?: string;
merged_tenant_user_id?: string;
role?: string;
}
/**
* Unified user picker hook — searches both channel_contacts and tenant_users.
* - Empty search: returns 30 most recent
* - With search: debounced server-side ILIKE search
* - Deduplicates merged contacts (shows tenant_user badge instead)
*/
/**
* @param source - Filter by source: "contact" | "tenant_user" | undefined (both).
* Use "tenant_user" for merge dialog / add tenant user (contacts excluded).
*/
export function useUserPicker(search: string, peerKind?: string, source?: "contact" | "tenant_user") {
const http = useHttp();
const [debouncedSearch, setDebouncedSearch] = useState("");
useEffect(() => {
const timer = setTimeout(() => setDebouncedSearch(search), 150);
return () => clearTimeout(timer);
}, [search]);
const { data, isLoading } = useQuery({
queryKey: queryKeys.users.search({ search: debouncedSearch, peerKind, source, limit: 30 }),
queryFn: async () => {
const params: Record<string, string> = { limit: "30" };
if (debouncedSearch) params.q = debouncedSearch;
if (peerKind) params.peer_kind = peerKind;
if (source) params.source = source;
const res = await http.get<{ results: UserPickerItem[] }>("/v1/users/search", params);
return res.results ?? [];
},
// Always enabled — empty search returns recent contacts
staleTime: 30_000,
});
const results = data ?? [];
/** Format results as ComboboxOptions with source badges. */
const options: ComboboxOption[] = useMemo(() =>
results.map((r) => {
const parts: string[] = [];
if (r.display_name) parts.push(r.display_name);
if (r.username) parts.push(`@${r.username}`);
parts.push(`(${r.id})`);
if (r.source === "contact" && r.channel_type) parts.push(`[${r.channel_type}]`);
if (r.source === "tenant_user") parts.push("[tenant]");
if (r.merged_tenant_user_id) parts.push(`→ ${r.merged_tenant_user_id}`);
return { value: r.id, label: parts.join(" ") };
}),
[results],
);
return { results, options, loading: isLoading };
}
@@ -56,5 +56,25 @@
"updateFailed": "Failed to update credential",
"deleted": "CLI credential deleted",
"deleteFailed": "Failed to delete credential"
},
"userCredentials": {
"title": "User Credentials",
"description": "Per-user environment variable overrides for {{name}}",
"userId": "User ID",
"userIdPlaceholder": "user-id or email",
"env": "Environment Variables",
"addEnv": "Add Variable",
"add": "Add",
"save": "Save",
"edit": "Edit",
"delete": "Delete",
"back": "Back",
"close": "Close",
"empty": "No user credentials configured",
"saved": "User credentials saved",
"saveFailed": "Failed to save user credentials",
"deleted": "User credentials deleted",
"deleteFailed": "Failed to delete user credentials",
"envRequired": "At least one environment variable is required"
}
}
@@ -56,5 +56,25 @@
"updateFailed": "Không thể cập nhật thông tin",
"deleted": "Đã xóa thông tin CLI",
"deleteFailed": "Không thể xóa thông tin"
},
"userCredentials": {
"title": "Thông tin người dùng",
"description": "Ghi đè biến môi trường cho từng người dùng của {{name}}",
"userId": "ID người dùng",
"userIdPlaceholder": "user-id hoặc email",
"env": "Biến môi trường",
"addEnv": "Thêm biến",
"add": "Thêm",
"save": "Lưu",
"edit": "Sửa",
"delete": "Xóa",
"back": "Quay lại",
"close": "Đóng",
"empty": "Chưa có thông tin người dùng nào",
"saved": "Đã lưu thông tin người dùng",
"saveFailed": "Không thể lưu thông tin người dùng",
"deleted": "Đã xóa thông tin người dùng",
"deleteFailed": "Không thể xóa thông tin người dùng",
"envRequired": "Cần ít nhất một biến môi trường"
}
}
@@ -56,5 +56,25 @@
"updateFailed": "更新凭证失败",
"deleted": "CLI 凭证已删除",
"deleteFailed": "删除凭证失败"
},
"userCredentials": {
"title": "用户凭证",
"description": "{{name}} 的每用户环境变量覆盖",
"userId": "用户 ID",
"userIdPlaceholder": "用户 ID 或邮箱",
"env": "环境变量",
"addEnv": "添加变量",
"add": "添加",
"save": "保存",
"edit": "编辑",
"delete": "删除",
"back": "返回",
"close": "关闭",
"empty": "未配置用户凭证",
"saved": "用户凭证已保存",
"saveFailed": "保存用户凭证失败",
"deleted": "用户凭证已删除",
"deleteFailed": "删除用户凭证失败",
"envRequired": "至少需要一个环境变量"
}
}
+4
View File
@@ -78,6 +78,10 @@ export const queryKeys = {
tenantUsers: {
all: ["tenantUsers"] as const,
},
users: {
all: ["users"] as const,
search: (params: Record<string, unknown>) => ["users", "search", params] as const,
},
tenants: {
all: ["tenants"] as const,
detail: (tenantId: string) => ["tenants", tenantId] as const,
@@ -1,6 +1,5 @@
import { useState, useEffect, useMemo, useRef, useLayoutEffect } from "react";
import { createPortal } from "react-dom";
import { Save, Loader2, Users, FileText, Search, UserPlus } from "lucide-react";
import { useState, useEffect, useMemo } from "react";
import { Save, Loader2, Users, FileText } from "lucide-react";
import { toast } from "@/stores/use-toast-store";
import { userFriendlyError } from "@/lib/error-utils";
import { useTranslation } from "react-i18next";
@@ -9,7 +8,7 @@ import { Badge } from "@/components/ui/badge";
import { Textarea } from "@/components/ui/textarea";
import { useContactResolver } from "@/hooks/use-contact-resolver";
import { useAgentInstances, type UserInstance } from "../hooks/use-agent-instances";
import { useContactSearch } from "../hooks/use-contact-search";
import { UserPickerCombobox } from "@/components/shared/user-picker-combobox";
interface AgentInstancesTabProps {
agentId: string;
@@ -60,8 +59,13 @@ export function AgentInstancesTab({ agentId }: AgentInstancesTabProps) {
// Existing instance user_ids for deduplication
const existingIDs = useMemo(() => new Set(instances.map((i) => i.user_id)), [instances]);
const handleContactSelect = (senderID: string) => {
setSelected(senderID);
const [addUserId, setAddUserId] = useState("");
const handleAddUser = (val: string) => {
setAddUserId(val);
if (val && !existingIDs.has(val)) {
setSelected(val);
}
};
if (loading) {
@@ -72,10 +76,14 @@ export function AgentInstancesTab({ agentId }: AgentInstancesTabProps) {
<div className="flex gap-4" style={{ minHeight: 400 }}>
{/* Instance list */}
<div className="w-64 shrink-0 space-y-1 overflow-y-auto rounded-md border p-2">
<ContactSearchBox
existingIDs={existingIDs}
onSelect={handleContactSelect}
/>
<div className="px-1 pb-2">
<UserPickerCombobox
value={addUserId}
onChange={handleAddUser}
placeholder={t("instances.searchContacts")}
className="w-full"
/>
</div>
{instances.length > 0 && (
<div className="px-2 pb-1 pt-1 text-xs font-medium text-muted-foreground">
{instances.length} instance{instances.length !== 1 ? "s" : ""}
@@ -135,98 +143,6 @@ export function AgentInstancesTab({ agentId }: AgentInstancesTabProps) {
);
}
/** Inline contact search dropdown for adding new instances. */
function ContactSearchBox({ existingIDs, onSelect }: { existingIDs: Set<string>; onSelect: (id: string) => void }) {
const { t } = useTranslation("agents");
const [search, setSearch] = useState("");
const [open, setOpen] = useState(false);
const { contacts } = useContactSearch(search);
const containerRef = useRef<HTMLDivElement>(null);
const dropdownRef = useRef<HTMLDivElement>(null);
const [dropdownStyle, setDropdownStyle] = useState<React.CSSProperties>({});
// Filter out contacts already in instances
const filtered = contacts.filter((c) => !existingIDs.has(c.sender_id));
// Compute dropdown position for portal rendering
useLayoutEffect(() => {
if (!open || !containerRef.current) return;
const rect = containerRef.current.getBoundingClientRect();
setDropdownStyle({
position: "fixed",
top: rect.bottom + 4,
left: rect.left,
width: rect.width,
zIndex: 9999,
});
}, [open, search]);
// Close on outside click
useEffect(() => {
if (!open) return;
const handler = (e: MouseEvent) => {
const target = e.target as Node;
if (
containerRef.current && !containerRef.current.contains(target) &&
(!dropdownRef.current || !dropdownRef.current.contains(target))
) {
setOpen(false);
}
};
document.addEventListener("mousedown", handler);
return () => document.removeEventListener("mousedown", handler);
}, [open]);
return (
<div ref={containerRef} className="relative px-1 pb-2">
<div className="relative">
<Search className="absolute left-2 top-1/2 h-3.5 w-3.5 -translate-y-1/2 text-muted-foreground" />
<input
value={search}
onChange={(e) => { setSearch(e.target.value); setOpen(true); }}
onFocus={() => search.length >= 2 && setOpen(true)}
placeholder={t("instances.searchContacts")}
className="h-8 w-full rounded-md border bg-transparent pl-7 pr-2 text-base md:text-xs placeholder:text-muted-foreground focus:outline-none focus:ring-1 focus:ring-ring"
/>
</div>
{open && search.length >= 2 && filtered.length > 0 && createPortal(
<div ref={dropdownRef} style={dropdownStyle} className="max-h-48 overflow-y-auto rounded-md border bg-popover p-1 shadow-md">
{filtered.map((c) => (
<button
key={c.id}
type="button"
onMouseDown={(e) => e.preventDefault()}
onClick={() => {
onSelect(c.sender_id);
setSearch("");
setOpen(false);
}}
className="flex w-full items-center gap-2 rounded-sm px-2 py-1.5 text-left text-xs hover:bg-accent hover:text-accent-foreground"
>
<UserPlus className="h-3.5 w-3.5 shrink-0 text-muted-foreground" />
<div className="min-w-0 flex-1">
<div className="truncate font-medium">
{c.display_name || c.sender_id}
</div>
<div className="flex items-center gap-1 text-[10px] text-muted-foreground">
{c.username && <span>@{c.username}</span>}
<Badge variant="outline" className="text-[9px] px-1 py-0">{c.channel_type}</Badge>
</div>
</div>
</button>
))}
</div>,
document.body,
)}
{open && search.length >= 2 && filtered.length === 0 && contacts.length === 0 && createPortal(
<div ref={dropdownRef} style={dropdownStyle} className="rounded-md border bg-popover p-3 text-center text-xs text-muted-foreground shadow-md">
{t("instances.noContactsFound")}
</div>,
document.body,
)}
</div>
);
}
function InstanceRow({ instance, isSelected, onClick, resolve }: { instance: UserInstance; isSelected: boolean; onClick: () => void; resolve: (id: string) => import("@/types/contact").ChannelContact | null }) {
const lastSeen = instance.last_seen_at ? formatRelative(instance.last_seen_at) : null;
@@ -8,8 +8,7 @@ import {
} from "@/components/ui/select";
import { Combobox, type ComboboxOption } from "@/components/ui/combobox";
import { useConfigPermissions, type ConfigPermission } from "../hooks/use-config-permissions";
import { useContactSearch } from "../hooks/use-contact-search";
import type { ChannelContact } from "@/types/contact";
import { UserPickerCombobox } from "@/components/shared/user-picker-combobox";
const CONFIG_TYPES = [
{ value: "file_writer", label: "File Writer", descKey: "permissions.types.file_writer_desc" },
@@ -48,19 +47,6 @@ export function AgentPermissionsTab({ agentId }: AgentPermissionsTabProps) {
const [scope, setScope] = useState("group:*");
const [permission, setPermission] = useState("allow");
const [adding, setAdding] = useState(false);
const [selectedContact, setSelectedContact] = useState<ChannelContact | null>(null);
const { contacts } = useContactSearch(userId);
const contactOptions: ComboboxOption[] = useMemo(() =>
contacts.map((c) => {
const name = c.display_name || c.sender_id;
const username = c.username ? ` @${c.username}` : "";
const channel = c.channel_type ? ` [${c.channel_type}]` : "";
return { value: c.sender_id, label: `${name}${username} (${c.sender_id})${channel}` };
}),
[contacts],
);
// Collect existing file_writer scopes for dynamic scope options
const existingFileWriterScopes = useMemo(() =>
@@ -88,25 +74,11 @@ export function AgentPermissionsTab({ agentId }: AgentPermissionsTabProps) {
useEffect(() => { load(); }, [load]);
const handleUserChange = (val: string) => {
setUserId(val);
const contact = contacts.find((c) => c.sender_id === val);
setSelectedContact(contact ?? null);
};
const handleAdd = async () => {
if (!userId.trim()) return;
setAdding(true);
const meta =
configType === "file_writer" && selectedContact
? {
displayName: selectedContact.display_name ?? "",
username: selectedContact.username ?? "",
}
: undefined;
await grant(scope, configType, userId.trim(), permission, meta);
await grant(scope, configType, userId.trim(), permission);
setUserId("");
setSelectedContact(null);
setAdding(false);
};
@@ -158,10 +130,9 @@ export function AgentPermissionsTab({ agentId }: AgentPermissionsTabProps) {
{/* Add Rule form */}
<div className="space-y-2">
<div className="flex flex-wrap items-end gap-2">
<Combobox
<UserPickerCombobox
value={userId}
onChange={handleUserChange}
options={contactOptions}
onChange={setUserId}
placeholder={t("permissions.userIdPlaceholder")}
className="flex-1 min-w-[160px]"
/>
@@ -1,13 +1,12 @@
import { useState, useMemo } from "react";
import { useState } from "react";
import { useTranslation } from "react-i18next";
import { Switch } from "@/components/ui/switch";
import { Badge } from "@/components/ui/badge";
import { Alert, AlertDescription } from "@/components/ui/alert";
import { Combobox } from "@/components/ui/combobox";
import { Shield, X, AlertTriangle, Plus, Brain } from "lucide-react";
import { Button } from "@/components/ui/button";
import type { WorkspaceSharingConfig } from "@/types/agent";
import { useContacts } from "@/pages/contacts/hooks/use-contacts";
import { UserPickerCombobox } from "@/components/shared/user-picker-combobox";
import { InfoLabel } from "./config-section";
const MAX_SHARED_USERS = 100;
@@ -22,22 +21,7 @@ export function WorkspaceSharingSection({ value, onChange }: WorkspaceSharingSec
const s = "configSections.workspaceSharing";
const [contactSearch, setContactSearch] = useState("");
// Fetch contacts for the combobox
const { contacts } = useContacts({ search: contactSearch, limit: 20 });
// Build combobox options from contacts, excluding already-added users
const existingUsers = value.shared_users ?? [];
const contactOptions = useMemo(() => {
const existing = new Set(existingUsers);
return contacts
.filter((c) => c.user_id && !existing.has(c.user_id))
.map((c) => ({
value: c.user_id!,
label: c.display_name
? `${c.display_name} (${c.user_id})`
: c.user_id!,
}));
}, [contacts, existingUsers]);
const addUser = (userId: string) => {
const trimmed = userId.trim();
@@ -139,10 +123,9 @@ export function WorkspaceSharingSection({ value, onChange }: WorkspaceSharingSec
</div>
)}
<div className="flex gap-2">
<Combobox
<UserPickerCombobox
value={contactSearch}
onChange={(val) => setContactSearch(val)}
options={contactOptions}
placeholder={t(`${s}.userIdPlaceholder`)}
className="flex-1"
/>
@@ -1,17 +1,16 @@
import { useState, useEffect, useRef, useLayoutEffect } from "react";
import { createPortal } from "react-dom";
import { Search, UserPlus } from "lucide-react";
import { useState } from "react";
import { useTranslation } from "react-i18next";
import { Badge } from "@/components/ui/badge";
import { useContactSearch } from "../../hooks/use-contact-search";
import type { ChannelContact } from "@/types/contact";
import { UserPickerCombobox } from "@/components/shared/user-picker-combobox";
import { useUserPicker } from "@/hooks/use-user-picker";
import type { UserPickerItem } from "@/hooks/use-user-picker";
/** Format a contact into a human-readable snippet for insertion into context files. */
function formatContactSnippet(c: ChannelContact): string {
/** Format a UserPickerItem into a human-readable snippet for insertion into context files. */
function formatSnippet(item: UserPickerItem): string {
const parts: string[] = [];
if (c.display_name) parts.push(c.display_name);
if (c.username) parts.push(`@${c.username}`);
parts.push(`${c.channel_type}:${c.sender_id}`);
if (item.display_name) parts.push(item.display_name);
if (item.username) parts.push(`@${item.username}`);
if (item.channel_type) parts.push(`${item.channel_type}:${item.id}`);
else parts.push(item.id);
return `- ${parts.join(" — ")}`;
}
@@ -22,90 +21,24 @@ interface ContactInsertSearchProps {
export function ContactInsertSearch({ onInsert }: ContactInsertSearchProps) {
const { t } = useTranslation("agents");
const [search, setSearch] = useState("");
const [open, setOpen] = useState(false);
const { contacts } = useContactSearch(search);
const containerRef = useRef<HTMLDivElement>(null);
const dropdownRef = useRef<HTMLDivElement>(null);
const [dropdownStyle, setDropdownStyle] = useState<React.CSSProperties>({});
const { results } = useUserPicker(search);
// Compute dropdown position for portal rendering
useLayoutEffect(() => {
if (!open || !containerRef.current) return;
const rect = containerRef.current.getBoundingClientRect();
setDropdownStyle({
position: "fixed",
top: rect.bottom + 4,
left: rect.left,
width: Math.min(rect.width, 384), // max-w-sm = 24rem = 384px
zIndex: 9999,
});
}, [open, search]);
// Close on outside click
useEffect(() => {
if (!open) return;
const handler = (e: MouseEvent) => {
const target = e.target as Node;
if (
containerRef.current && !containerRef.current.contains(target) &&
(!dropdownRef.current || !dropdownRef.current.contains(target))
) {
setOpen(false);
}
};
document.addEventListener("mousedown", handler);
return () => document.removeEventListener("mousedown", handler);
}, [open]);
const handleSelect = (c: ChannelContact) => {
onInsert(formatContactSnippet(c));
setSearch("");
setOpen(false);
const handleChange = (val: string) => {
const item = results.find((r) => r.id === val);
if (item) {
onInsert(formatSnippet(item));
setSearch("");
} else {
setSearch(val);
}
};
return (
<div ref={containerRef} className="relative">
<div className="relative">
<Search className="absolute left-2 top-1/2 h-3.5 w-3.5 -translate-y-1/2 text-muted-foreground" />
<input
value={search}
onChange={(e) => { setSearch(e.target.value); setOpen(true); }}
onFocus={() => search.length >= 2 && setOpen(true)}
placeholder={t("files.insertContact")}
className="h-8 w-full max-w-sm rounded-md border bg-transparent pl-7 pr-2 text-base md:text-xs placeholder:text-muted-foreground focus:outline-none focus:ring-1 focus:ring-ring"
/>
</div>
{open && search.length >= 2 && contacts.length > 0 && createPortal(
<div ref={dropdownRef} style={dropdownStyle} className="max-h-48 overflow-y-auto rounded-md border bg-popover p-1 shadow-md">
{contacts.map((c) => (
<button
key={c.id}
type="button"
onMouseDown={(e) => e.preventDefault()}
onClick={() => handleSelect(c)}
className="flex w-full items-center gap-2 rounded-sm px-2 py-1.5 text-left text-xs hover:bg-accent hover:text-accent-foreground"
>
<UserPlus className="h-3.5 w-3.5 shrink-0 text-muted-foreground" />
<div className="min-w-0 flex-1">
<div className="truncate font-medium">
{c.display_name || c.sender_id}
</div>
<div className="flex items-center gap-1 text-[10px] text-muted-foreground">
{c.username && <span>@{c.username}</span>}
<Badge variant="outline" className="text-[9px] px-1 py-0">{c.channel_type}</Badge>
</div>
</div>
</button>
))}
</div>,
document.body,
)}
{open && search.length >= 2 && contacts.length === 0 && createPortal(
<div ref={dropdownRef} style={dropdownStyle} className="rounded-md border bg-popover p-3 text-center text-xs text-muted-foreground shadow-md">
{t("instances.noContactsFound")}
</div>,
document.body,
)}
</div>
<UserPickerCombobox
value={search}
onChange={handleChange}
placeholder={t("files.insertContact")}
className="h-8 w-full max-w-sm"
/>
);
}
@@ -29,7 +29,6 @@ export function ChannelDetailPage({ instanceId, onBack, onDelete }: ChannelDetai
listManagers,
addManager,
removeManager,
listContacts,
} = useChannelDetail(instanceId);
const { agents } = useAgents();
const { channels } = useChannels();
@@ -107,7 +106,6 @@ export function ChannelDetailPage({ instanceId, onBack, onDelete }: ChannelDetai
listManagers={listManagers}
addManager={addManager}
removeManager={removeManager}
listContacts={listContacts}
/>
</TabsContent>
</Tabs>
@@ -4,18 +4,15 @@ import { Plus, Trash2, Loader2, RefreshCw, Users, ChevronDown, ChevronRight } fr
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Label } from "@/components/ui/label";
import { Combobox } from "@/components/ui/combobox";
import { EmptyState } from "@/components/shared/empty-state";
import { useContactPicker } from "@/hooks/use-contact-picker";
import { UserPickerCombobox } from "@/components/shared/user-picker-combobox";
import type { GroupManagerGroupInfo, GroupManagerData } from "../hooks/use-channel-detail";
import type { ChannelContact } from "@/types/contact";
interface ChannelManagersTabProps {
listManagerGroups: () => Promise<GroupManagerGroupInfo[]>;
listManagers: (groupId: string) => Promise<GroupManagerData[]>;
addManager: (groupId: string, userId: string, displayName?: string, username?: string) => Promise<void>;
removeManager: (groupId: string, userId: string) => Promise<void>;
listContacts: (search: string, channelType?: string) => Promise<ChannelContact[]>;
}
/** Strips the "group:<channel>:" prefix for display, e.g. "group:telegram:-100123" → "-100123" */
@@ -29,22 +26,15 @@ function shortGroupId(id: string): string {
interface InlineAddFormProps {
groupId?: string;
showGroupField?: boolean;
listContacts: (search: string) => Promise<ChannelContact[]>;
onAdd: (groupId: string, userId: string, displayName: string, username: string) => Promise<void>;
}
function InlineAddForm({ groupId, showGroupField, listContacts, onAdd }: InlineAddFormProps) {
function InlineAddForm({ groupId, showGroupField, onAdd }: InlineAddFormProps) {
const { t } = useTranslation("channels");
const [formGroupId, setFormGroupId] = useState("");
const [userId, setUserId] = useState("");
const [adding, setAdding] = useState(false);
const [error, setError] = useState("");
const { options: contactOptions, searchContacts, getContact, clearOptions } = useContactPicker(listContacts);
const handleUserIdChange = (val: string) => {
setUserId(val);
searchContacts(val);
};
const handleSubmit = async () => {
const gid = groupId || formGroupId.trim();
@@ -56,13 +46,8 @@ function InlineAddForm({ groupId, showGroupField, listContacts, onAdd }: InlineA
setAdding(true);
setError("");
try {
// Auto-fill display name and username from selected contact
const contact = getContact(uid);
const displayName = contact?.display_name ?? "";
const username = contact?.username ?? "";
await onAdd(gid, uid, displayName, username);
await onAdd(gid, uid, "", "");
setUserId("");
clearOptions();
if (!groupId) setFormGroupId("");
} catch (err) {
setError(err instanceof Error ? err.message : t("detail.managers.addForm.errors.failedAdd"));
@@ -88,10 +73,9 @@ function InlineAddForm({ groupId, showGroupField, listContacts, onAdd }: InlineA
</div>
<div className="grid gap-1.5 flex-1 min-w-[180px]">
<Label className="text-xs">{t("detail.managers.addForm.userId")}</Label>
<Combobox
<UserPickerCombobox
value={userId}
onChange={handleUserIdChange}
options={contactOptions}
onChange={setUserId}
placeholder={t("detail.managers.addForm.userIdPlaceholder")}
/>
</div>
@@ -114,10 +98,9 @@ function InlineAddForm({ groupId, showGroupField, listContacts, onAdd }: InlineA
<div className="flex items-end gap-2">
<div className="grid gap-1 flex-1 min-w-[140px]">
<Label className="text-xs text-muted-foreground">{t("detail.managers.addForm.userId")}</Label>
<Combobox
<UserPickerCombobox
value={userId}
onChange={handleUserIdChange}
options={contactOptions}
onChange={setUserId}
placeholder={t("detail.managers.addForm.userIdPlaceholder")}
className="h-8"
/>
@@ -143,7 +126,6 @@ export function ChannelManagersTab({
listManagers,
addManager,
removeManager,
listContacts,
}: ChannelManagersTabProps) {
const { t } = useTranslation("channels");
const [groups, setGroups] = useState<GroupManagerGroupInfo[]>([]);
@@ -318,7 +300,6 @@ export function ChannelManagersTab({
<InlineAddForm
groupId={g.group_id}
listContacts={listContacts}
onAdd={handleAddManager}
/>
</div>
@@ -333,7 +314,6 @@ export function ChannelManagersTab({
{/* Add to new group */}
<InlineAddForm
showGroupField
listContacts={listContacts}
onAdd={handleAddManager}
/>
</div>
@@ -1,6 +1,6 @@
import { useState } from "react";
import { useTranslation } from "react-i18next";
import { KeyRound, Plus, RefreshCw, Pencil, Trash2 } from "lucide-react";
import { KeyRound, Plus, RefreshCw, Pencil, Trash2, Users } from "lucide-react";
import { Button } from "@/components/ui/button";
import { Badge } from "@/components/ui/badge";
import { PageHeader } from "@/components/shared/page-header";
@@ -11,6 +11,7 @@ import { useMinLoading } from "@/hooks/use-min-loading";
import { useDeferredLoading } from "@/hooks/use-deferred-loading";
import { useCliCredentials, useCliCredentialPresets } from "./hooks/use-cli-credentials";
import { CliCredentialFormDialog } from "./cli-credential-form-dialog";
import { CLIUserCredentialsDialog } from "./cli-user-credentials-dialog";
import type { SecureCLIBinary, CLICredentialInput } from "./hooks/use-cli-credentials";
export function CliCredentialsPage() {
@@ -21,6 +22,7 @@ export function CliCredentialsPage() {
const [editItem, setEditItem] = useState<SecureCLIBinary | null>(null);
const [deleteTarget, setDeleteTarget] = useState<SecureCLIBinary | null>(null);
const [deleteLoading, setDeleteLoading] = useState(false);
const [userCredsTarget, setUserCredsTarget] = useState<SecureCLIBinary | null>(null);
const { items, loading, refresh, createCredential, updateCredential, deleteCredential } =
useCliCredentials();
@@ -128,6 +130,9 @@ export function CliCredentialsPage() {
<td className="px-4 py-3 text-muted-foreground">{item.timeout_seconds}s</td>
<td className="px-4 py-3 text-right">
<div className="flex items-center justify-end gap-1">
<Button variant="ghost" size="sm" onClick={() => setUserCredsTarget(item)} title={t("userCredentials.title")}>
<Users className="h-3.5 w-3.5" />
</Button>
<Button
variant="ghost"
size="sm"
@@ -172,6 +177,14 @@ export function CliCredentialsPage() {
onConfirm={handleDelete}
loading={deleteLoading}
/>
{userCredsTarget && (
<CLIUserCredentialsDialog
open={!!userCredsTarget}
onOpenChange={(open: boolean) => !open && setUserCredsTarget(null)}
binary={userCredsTarget}
/>
)}
</div>
);
}
@@ -0,0 +1,265 @@
import { useState, useEffect, useCallback } from "react";
import { useTranslation } from "react-i18next";
import { KeyRound, Loader2, Plus, Trash2, Pencil } from "lucide-react";
import {
Dialog,
DialogContent,
DialogHeader,
DialogTitle,
DialogDescription,
DialogFooter,
} from "@/components/ui/dialog";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Label } from "@/components/ui/label";
import { Badge } from "@/components/ui/badge";
import { KeyValueEditor } from "@/components/shared/key-value-editor";
import { toast } from "@/stores/use-toast-store";
import { useHttp } from "@/hooks/use-ws";
import i18next from "i18next";
import type { SecureCLIBinary } from "./hooks/use-cli-credentials";
interface UserCredEntry {
id: string;
binary_id: string;
user_id: string;
has_env: boolean;
created_at: string;
updated_at: string;
}
interface CLIUserCredentialsDialogProps {
open: boolean;
onOpenChange: (open: boolean) => void;
binary: SecureCLIBinary;
}
const SENSITIVE_ENV_RE = /^.*(key|secret|token|password|credential).*$/i;
const isSensitiveEnv = (key: string) => SENSITIVE_ENV_RE.test(key.trim());
type ViewState = "list" | "form";
export function CLIUserCredentialsDialog({ open, onOpenChange, binary }: CLIUserCredentialsDialogProps) {
const { t } = useTranslation("cli-credentials");
const http = useHttp();
const [view, setView] = useState<ViewState>("list");
const [entries, setEntries] = useState<UserCredEntry[]>([]);
const [loadingList, setLoadingList] = useState(false);
// Form state
const [editEntry, setEditEntry] = useState<UserCredEntry | null>(null);
const [userId, setUserId] = useState("");
const [env, setEnv] = useState<Record<string, string>>({});
const [saving, setSaving] = useState(false);
const [deleting, setDeletingId] = useState<string | null>(null);
const loadList = useCallback(async () => {
setLoadingList(true);
try {
const res = await http.get<{ user_credentials: UserCredEntry[] }>(
`/v1/cli-credentials/${binary.id}/user-credentials`,
);
setEntries(res.user_credentials ?? []);
} catch {
// silently ignore
} finally {
setLoadingList(false);
}
}, [http, binary.id]);
useEffect(() => {
if (!open) return;
setView("list");
setEditEntry(null);
setUserId("");
setEnv({});
loadList();
}, [open, loadList]);
const openAdd = () => {
setEditEntry(null);
setUserId("");
setEnv({});
setView("form");
};
const openEdit = async (entry: UserCredEntry) => {
setEditEntry(entry);
setUserId(entry.user_id);
setEnv({});
setView("form");
// Load existing env for edit
try {
const res = await http.get<{ user_id: string; env: Record<string, string> | null }>(
`/v1/cli-credentials/${binary.id}/user-credentials/${entry.user_id}`,
);
setEnv(res.env ?? {});
} catch {
// leave env empty — user can re-enter
}
};
const handleSave = async () => {
const uid = userId.trim();
if (!uid) return;
if (Object.keys(env).length === 0) {
toast.error(i18next.t("cli-credentials:userCredentials.envRequired"));
return;
}
setSaving(true);
try {
await http.put(`/v1/cli-credentials/${binary.id}/user-credentials/${uid}`, { env });
toast.success(i18next.t("cli-credentials:userCredentials.saved"));
await loadList();
setView("list");
} catch (err) {
toast.error(
i18next.t("cli-credentials:userCredentials.saveFailed"),
err instanceof Error ? err.message : "",
);
} finally {
setSaving(false);
}
};
const handleDelete = async (entry: UserCredEntry) => {
setDeletingId(entry.id);
try {
await http.delete(`/v1/cli-credentials/${binary.id}/user-credentials/${entry.user_id}`);
toast.success(i18next.t("cli-credentials:userCredentials.deleted"));
await loadList();
} catch (err) {
toast.error(
i18next.t("cli-credentials:userCredentials.deleteFailed"),
err instanceof Error ? err.message : "",
);
} finally {
setDeletingId(null);
}
};
return (
<Dialog open={open} onOpenChange={onOpenChange}>
<DialogContent className="sm:max-w-lg">
<DialogHeader>
<DialogTitle className="flex items-center gap-2">
<KeyRound className="h-4 w-4" />
{t("userCredentials.title")}
</DialogTitle>
<DialogDescription>
{t("userCredentials.description", { name: binary.binary_name })}
</DialogDescription>
</DialogHeader>
{view === "list" ? (
<>
{loadingList ? (
<div className="flex justify-center py-8">
<Loader2 className="h-5 w-5 animate-spin text-muted-foreground" />
</div>
) : entries.length === 0 ? (
<p className="py-6 text-center text-sm text-muted-foreground">
{t("userCredentials.empty")}
</p>
) : (
<div className="flex flex-col gap-2 max-h-[50vh] overflow-y-auto pr-1">
{entries.map((entry) => (
<div
key={entry.id}
className="flex items-center justify-between rounded-md border px-3 py-2"
>
<div className="flex items-center gap-2 min-w-0">
<span className="font-mono text-sm truncate">{entry.user_id}</span>
{entry.has_env && (
<Badge variant="secondary" className="shrink-0 text-xs">
env
</Badge>
)}
</div>
<div className="flex items-center gap-1 shrink-0">
<Button
variant="ghost"
size="icon"
className="h-8 w-8"
onClick={() => openEdit(entry)}
title={t("userCredentials.edit")}
>
<Pencil className="h-3.5 w-3.5" />
</Button>
<Button
variant="ghost"
size="icon"
className="h-8 w-8 text-destructive hover:text-destructive"
onClick={() => handleDelete(entry)}
disabled={deleting === entry.id}
title={t("userCredentials.delete")}
>
{deleting === entry.id ? (
<Loader2 className="h-3.5 w-3.5 animate-spin" />
) : (
<Trash2 className="h-3.5 w-3.5" />
)}
</Button>
</div>
</div>
))}
</div>
)}
<DialogFooter>
<Button variant="outline" onClick={() => onOpenChange(false)}>
{t("userCredentials.close")}
</Button>
<Button onClick={openAdd} className="gap-1">
<Plus className="h-3.5 w-3.5" />
{t("userCredentials.add")}
</Button>
</DialogFooter>
</>
) : (
<>
<div className="flex flex-col gap-4 max-h-[60vh] overflow-y-auto pr-1">
<div className="flex flex-col gap-1.5">
<Label htmlFor="uc-user-id">{t("userCredentials.userId")}</Label>
<Input
id="uc-user-id"
value={userId}
onChange={(e) => setUserId(e.target.value)}
placeholder={t("userCredentials.userIdPlaceholder")}
disabled={!!editEntry}
className="text-base md:text-sm font-mono"
/>
</div>
<div className="flex flex-col gap-1.5">
<Label>{t("userCredentials.env")}</Label>
<KeyValueEditor
value={env}
onChange={setEnv}
keyPlaceholder="ENV_KEY"
valuePlaceholder="value"
addLabel={t("userCredentials.addEnv")}
maskValue={isSensitiveEnv}
/>
</div>
</div>
<DialogFooter>
<Button
variant="outline"
onClick={() => setView("list")}
disabled={saving}
>
{t("userCredentials.back")}
</Button>
<Button onClick={handleSave} disabled={saving}>
{saving ? <Loader2 className="h-3.5 w-3.5 animate-spin mr-1" /> : null}
{t("userCredentials.save")}
</Button>
</DialogFooter>
</>
)}
</DialogContent>
</Dialog>
);
}
@@ -1,10 +1,9 @@
import { useEffect, useState } from "react";
import { useTranslation } from "react-i18next";
import { Construction, Merge } from "lucide-react";
import { Merge } from "lucide-react";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Label } from "@/components/ui/label";
import { Combobox } from "@/components/ui/combobox";
import {
Dialog,
DialogContent,
@@ -16,7 +15,7 @@ import {
import { toast } from "@/stores/use-toast-store";
import type { ChannelContact } from "@/types/contact";
import { useContactMerge } from "./hooks/use-contact-merge";
import { useTenantUsersList } from "./hooks/use-tenant-users-list";
import { UserPickerCombobox } from "@/components/shared/user-picker-combobox";
interface MergeContactsDialogProps {
open: boolean;
@@ -35,15 +34,12 @@ export function MergeContactsDialog({
}: MergeContactsDialogProps) {
const { t } = useTranslation("contacts");
const { merge } = useContactMerge();
const { users } = useTenantUsersList();
const [mode, setMode] = useState<MergeMode>("existing");
const [selectedUserId, setSelectedUserId] = useState("");
const [newDisplayName, setNewDisplayName] = useState("");
const [newUserId, setNewUserId] = useState("");
// TODO: re-enable when merge feature is ready
// const [submitting, setSubmitting] = useState(false);
const setSubmitting = (_v: boolean) => {};
const [submitting, setSubmitting] = useState(false);
// Reset form state when dialog opens
useEffect(() => {
@@ -59,11 +55,6 @@ export function MergeContactsDialog({
const defaultUserId =
selectedContacts[0]?.username || selectedContacts[0]?.sender_id || "";
const userOptions = users.map((u) => ({
value: u.id,
label: u.display_name || u.user_id,
}));
const handleSubmit = async () => {
const contactIds = selectedContacts.map((c) => c.id);
setSubmitting(true);
@@ -92,8 +83,7 @@ export function MergeContactsDialog({
}
};
// TODO: re-enable when merge feature is ready
// const canSubmit = mode === "existing" ? !!selectedUserId : !!(newUserId || defaultUserId);
const canSubmit = mode === "existing" ? !!selectedUserId : !!(newUserId || defaultUserId);
return (
<Dialog open={open} onOpenChange={onOpenChange}>
@@ -106,13 +96,7 @@ export function MergeContactsDialog({
<DialogDescription>{t("merge.dialogDescription")}</DialogDescription>
</DialogHeader>
{/* Coming soon banner */}
<div className="flex items-center gap-2 rounded-md border border-amber-500/30 bg-amber-500/10 px-3 py-2 text-sm text-amber-700 dark:text-amber-400">
<Construction className="h-4 w-4 shrink-0" />
<span className="font-medium">{t("merge.comingSoon")}</span>
</div>
<div className="space-y-4 py-2 pointer-events-none opacity-50">
<div className="space-y-4 py-2">
{/* Mode selection — simple radio buttons */}
<div className="space-y-2">
<label className="flex items-center gap-2 cursor-pointer">
@@ -128,11 +112,11 @@ export function MergeContactsDialog({
{mode === "existing" && (
<div className="ml-6">
<Combobox
<UserPickerCombobox
value={selectedUserId}
onChange={setSelectedUserId}
options={userOptions}
placeholder={t("merge.selectUser")}
source="tenant_user"
/>
</div>
)}
@@ -187,7 +171,7 @@ export function MergeContactsDialog({
<Button variant="outline" onClick={() => onOpenChange(false)}>
{t("merge.cancel", { defaultValue: "Cancel" })}
</Button>
<Button onClick={handleSubmit} disabled>
<Button onClick={handleSubmit} disabled={!canSubmit || submitting}>
{t("merge.confirm")}
</Button>
</DialogFooter>
@@ -1,4 +1,4 @@
import { useState, useEffect, useMemo } from "react";
import { useState, useEffect } from "react";
import { useTranslation } from "react-i18next";
import { KeyRound, Loader2 } from "lucide-react";
import {
@@ -14,11 +14,10 @@ import { Input } from "@/components/ui/input";
import { Label } from "@/components/ui/label";
import { Badge } from "@/components/ui/badge";
import { KeyValueEditor } from "@/components/shared/key-value-editor";
import { Combobox } from "@/components/ui/combobox";
import { UserPickerCombobox } from "@/components/shared/user-picker-combobox";
import { toast } from "@/stores/use-toast-store";
import { useAuthStore } from "@/stores/use-auth-store";
import { useTenants } from "@/hooks/use-tenants";
import { useTenantUsersList } from "@/pages/contacts/hooks/use-tenant-users-list";
import i18next from "i18next";
import type { MCPServerData, MCPUserCredentialStatus, MCPUserCredentialInput } from "./hooks/use-mcp";
@@ -51,7 +50,6 @@ export function MCPUserCredentialsDialog({
const role = useAuthStore((s) => s.role);
const currentUserId = useAuthStore((s) => s.userId);
const { currentTenant } = useTenants();
const { users } = useTenantUsersList();
const canManageUsers =
role === "admin" || role === "owner" ||
@@ -69,14 +67,6 @@ export function MCPUserCredentialsDialog({
const [headers, setHeaders] = useState<Record<string, string>>({});
const [env, setEnv] = useState<Record<string, string>>({});
const userOptions = useMemo(
() =>
users.map((u) => ({
value: u.user_id,
label: u.display_name || u.user_id,
})),
[users],
);
// Reset selected user when dialog opens
useEffect(() => {
@@ -152,12 +142,12 @@ export function MCPUserCredentialsDialog({
{canManageUsers && (
<div className="flex flex-col gap-1.5">
<Label>{t("userCredentials.selectUser")}</Label>
<Combobox
<UserPickerCombobox
value={selectedUserId}
onChange={(val) => setSelectedUserId(val)}
options={userOptions}
onChange={setSelectedUserId}
placeholder={t("userCredentials.selectUser")}
/>
source="tenant_user"
/>
</div>
)}
@@ -3,7 +3,6 @@ import { useParams, useNavigate } from "react-router";
import { useTranslation } from "react-i18next";
import { ArrowLeft, Plus, RefreshCw, Users, Trash2, Calendar, Hash, Shield } from "lucide-react";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Label } from "@/components/ui/label";
import { Badge } from "@/components/ui/badge";
import {
@@ -25,6 +24,7 @@ import {
import { PageHeader } from "@/components/shared/page-header";
import { EmptyState } from "@/components/shared/empty-state";
import { TableSkeleton } from "@/components/shared/loading-skeleton";
import { UserPickerCombobox } from "@/components/shared/user-picker-combobox";
import { useDeferredLoading } from "@/hooks/use-deferred-loading";
import { useMinLoading } from "@/hooks/use-min-loading";
import { useTenantDetail } from "./hooks/use-tenant-detail";
@@ -183,9 +183,14 @@ export function TenantDetailPage() {
</DialogHeader>
<div className="space-y-4 py-2">
<div className="space-y-1.5">
<Label htmlFor="add-user-id">{t("userId")}</Label>
<Input id="add-user-id" value={userId} onChange={(e) => setUserId(e.target.value)}
placeholder="user-id" className="text-base md:text-sm" />
<Label>{t("userId")}</Label>
<UserPickerCombobox
value={userId}
onChange={setUserId}
placeholder="user-id"
source="tenant_user"
allowCustom={true}
/>
</div>
<div className="space-y-1.5">
<Label>{t("selectRole")}</Label>