diff --git a/cli/exp_scaletest.go b/cli/exp_scaletest.go index a4d5b14d65..c49a228a54 100644 --- a/cli/exp_scaletest.go +++ b/cli/exp_scaletest.go @@ -396,6 +396,7 @@ type workspaceTargetFlags struct { template string targetWorkspaces string useHostLogin bool + allowEmpty bool } // attach adds the workspace target flags to the given options set. @@ -463,6 +464,9 @@ func (f *workspaceTargetFlags) getTargetedWorkspaces(ctx context.Context, client // Validate range if len(workspaces) == 0 { + if f.allowEmpty { + return nil, nil + } return nil, xerrors.Errorf("no scaletest workspaces exist") } if targetEnd > len(workspaces) { diff --git a/cli/exp_scaletest_chat.go b/cli/exp_scaletest_chat.go index bbde5f67ab..992a1944d9 100644 --- a/cli/exp_scaletest_chat.go +++ b/cli/exp_scaletest_chat.go @@ -11,8 +11,6 @@ import ( "github.com/prometheus/client_golang/prometheus/promhttp" "golang.org/x/xerrors" - "cdr.dev/slog/v3" - "cdr.dev/slog/v3/sloggers/sloghuman" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/scaletest/chat" "github.com/coder/coder/v2/scaletest/harness" @@ -22,17 +20,18 @@ import ( func (r *RootCmd) scaletestChat() *serpent.Command { var ( - chatsPerWorkspace int64 - prompt string - turns int64 - turnStartDelay time.Duration - llmMockURL string - targetFlags = &workspaceTargetFlags{} - tracingFlags = &scaletestTracingFlags{} - prometheusFlags = &scaletestPrometheusFlags{} - timeoutStrategy = &timeoutFlags{} - cleanupStrategy = newScaletestCleanupStrategy() - output = &scaletestOutputFlags{} + chatsPerWorkspace int64 + prompt string + turns int64 + turnStartDelay time.Duration + llmMockURL string + providerPropagationWait time.Duration + targetFlags = &workspaceTargetFlags{allowEmpty: true} + tracingFlags = &scaletestTracingFlags{} + prometheusFlags = &scaletestPrometheusFlags{} + timeoutStrategy = &timeoutFlags{} + cleanupStrategy = newScaletestCleanupStrategy() + output = &scaletestOutputFlags{} ) cmd := &serpent.Command{ @@ -72,8 +71,13 @@ func (r *RootCmd) scaletestChat() *serpent.Command { return err } - logger := slog.Make(sloghuman.Sink(inv.Stderr)).Leveled(slog.LevelDebug) - modelConfigID, err := chat.EnsureScaletestModelConfig(ctx, codersdk.NewExperimentalClient(client), logger, llmMockURL) + if len(workspaces) == 0 { + workspaces = append(workspaces, codersdk.Workspace{OrganizationID: me.OrganizationIDs[0]}) + _, _ = fmt.Fprintln(inv.Stderr, "No scaletest workspaces found; running chats without workspace context.") + } + + logger := inv.Logger + modelConfigID, err := chat.EnsureScaletestModelConfig(ctx, client, logger, llmMockURL, providerPropagationWait) if err != nil { return err } @@ -154,7 +158,7 @@ func (r *RootCmd) scaletestChat() *serpent.Command { // Run the chat harness in the background so the CLI can release the // follow-up turns after every runner finishes its initial turn. totalChats := int64(len(workspaces)) * chatsPerWorkspace - _, _ = fmt.Fprintf(inv.Stderr, "Starting chat scale test with %d chats across %d workspaces...\n", totalChats, len(workspaces)) + _, _ = fmt.Fprintf(inv.Stderr, "Starting chat scale test with %d chats across %d targets...\n", totalChats, len(workspaces)) testCtx, testCancel := timeoutStrategy.toContext(ctx) defer testCancel() testDone := make(chan error, 1) @@ -243,6 +247,13 @@ func (r *RootCmd) scaletestChat() *serpent.Command { Value: serpent.StringOf(&llmMockURL), Required: true, }, + { + Flag: "provider-propagation-wait", + Description: "Time to wait after creating or updating the mock LLM provider so every coderd replica's cached provider config expires. The default exceeds the server-side cache TTL.", + Default: chat.DefaultProviderPropagationWait.String(), + Value: serpent.DurationOf(&providerPropagationWait), + Hidden: true, + }, } targetFlags.attach(&cmd.Options) output.attach(&cmd.Options) diff --git a/cli/exp_scaletest_chat_test.go b/cli/exp_scaletest_chat_test.go new file mode 100644 index 0000000000..f5c2db8444 --- /dev/null +++ b/cli/exp_scaletest_chat_test.go @@ -0,0 +1,141 @@ +//go:build !slim + +package cli_test + +import ( + "bytes" + "context" + "io" + "strings" + "testing" + + "github.com/google/uuid" + "github.com/stretchr/testify/require" + + "cdr.dev/slog/v3" + "cdr.dev/slog/v3/sloggers/sloghuman" + "github.com/coder/coder/v2/cli/clitest" + "github.com/coder/coder/v2/coderd/coderdtest" + "github.com/coder/coder/v2/codersdk" + "github.com/coder/coder/v2/scaletest/llmmock" + "github.com/coder/coder/v2/testutil" +) + +const scaletestChatPrompt = "Reply with one short sentence from the scaletest." + +func TestScaleTestChat(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + values := coderdtest.DeploymentValues(t, func(dv *codersdk.DeploymentValues) { + require.NoError(t, dv.AI.BridgeConfig.Enabled.Set("true")) + // Keep AI Gateway routing disabled so the chat uses the direct model + // route to the mock provider, avoiding the need for an aibridged daemon. + require.NoError(t, dv.AI.Chat.AIGatewayRoutingEnabled.Set("false")) + }) + client := coderdtest.New(t, &coderdtest.Options{ + DeploymentValues: values, + }) + coderdtest.CreateFirstUser(t, client) + + server := new(llmmock.Server) + require.NoError(t, server.Start(context.Background(), llmmock.Config{ + Address: "127.0.0.1:0", + Logger: slog.Make(sloghuman.Sink(io.Discard)).Leveled(slog.LevelDebug), + })) + t.Cleanup(func() { + require.NoError(t, server.Stop()) + }) + mockURL := server.APIAddress() + "/v1" + + inv, root := clitest.New(t, + "exp", "scaletest", "chat", + "--chats-per-workspace", "1", + "--turns", "1", + "--prompt", scaletestChatPrompt, + "--timeout", "30s", + "--job-timeout", "30s", + "--cleanup-timeout", "30s", + "--cleanup-job-timeout", "30s", + "--scaletest-prometheus-address", "127.0.0.1:0", + "--scaletest-prometheus-wait", "0s", + "--provider-propagation-wait", "10ms", + "--llm-mock-url", mockURL, + ) + //nolint:gocritic // The scaletest chat command requires an admin client. + clitest.SetupConfig(t, client, root) + + var stderr bytes.Buffer + inv.Stdout = io.Discard + inv.Stderr = &stderr + + err := inv.WithContext(ctx).Run() + require.NoError(t, err, stderr.String()) + require.Contains(t, stderr.String(), "Scale test passed: 1/1 runs succeeded") + + provider, err := client.AIProvider(ctx, "coder-scaletest-mock") + require.NoError(t, err) + require.Equal(t, mockURL, provider.BaseURL) + + expClient := codersdk.NewExperimentalClient(client) + configs, err := expClient.ListChatModelConfigs(ctx) + require.NoError(t, err) + matchingConfigs := scaletestModelConfigsForProvider(configs, provider.ID) + require.Len(t, matchingConfigs, 1) + require.True(t, matchingConfigs[0].Enabled) + + chats, err := expClient.ListChats(ctx, &codersdk.ListChatsOptions{Query: "archived:true"}) + require.NoError(t, err) + + var scaletestMessages []codersdk.ChatMessage + for _, chat := range chats { + resp, err := expClient.GetChatMessages(ctx, chat.ID, nil) + require.NoError(t, err) + if userText, ok := chatMessageText(resp.Messages, codersdk.ChatMessageRoleUser); ok && + strings.Contains(userText, scaletestChatPrompt) { + scaletestMessages = resp.Messages + break + } + } + require.NotEmpty(t, scaletestMessages) + assistantText, ok := chatMessageText(scaletestMessages, codersdk.ChatMessageRoleAssistant) + require.True(t, ok, "expected an assistant reply in the scaletest chat") + require.NotEmpty(t, assistantText) +} + +// chatMessageText concatenates the text parts of every message with the given +// role, reporting whether any such message was found. It aggregates across +// messages because the API returns them newest-first and a turn can produce +// more than one message per role. +func chatMessageText(messages []codersdk.ChatMessage, role codersdk.ChatMessageRole) (string, bool) { + var ( + b strings.Builder + found bool + ) + for _, msg := range messages { + if msg.Role != role { + continue + } + found = true + for _, part := range msg.Content { + if part.Type == codersdk.ChatMessagePartTypeText { + _, _ = b.WriteString(part.Text) + } + } + } + return b.String(), found +} + +func scaletestModelConfigsForProvider(configs []codersdk.ChatModelConfig, providerID uuid.UUID) []codersdk.ChatModelConfig { + matches := make([]codersdk.ChatModelConfig, 0, 1) + for _, config := range configs { + if config.AIProviderID == nil || *config.AIProviderID != providerID { + continue + } + if config.Model != "scaletest-model" { + continue + } + matches = append(matches, config) + } + return matches +} diff --git a/scaletest/chat/client.go b/scaletest/chat/client.go index 552bbd87e1..bb2ad29c74 100644 --- a/scaletest/chat/client.go +++ b/scaletest/chat/client.go @@ -19,36 +19,11 @@ type chatClient interface { UpdateChat(ctx context.Context, chatID uuid.UUID, req codersdk.UpdateChatRequest) error } -type sdkChatClient struct { - client *codersdk.ExperimentalClient +var _ chatClient = (*codersdk.ExperimentalClient)(nil) + +type chatModelConfigClient interface { + ListChatModelConfigs(ctx context.Context) ([]codersdk.ChatModelConfig, error) + CreateChatModelConfig(ctx context.Context, req codersdk.CreateChatModelConfigRequest) (codersdk.ChatModelConfig, error) } -func newChatClient(client *codersdk.Client) chatClient { - return &sdkChatClient{client: codersdk.NewExperimentalClient(client)} -} - -func (c *sdkChatClient) SetLogger(logger slog.Logger) { - c.client.SetLogger(logger) -} - -func (c *sdkChatClient) SetLogBodies(logBodies bool) { - c.client.SetLogBodies(logBodies) -} - -func (c *sdkChatClient) CreateChat(ctx context.Context, req codersdk.CreateChatRequest) (codersdk.Chat, error) { - return c.client.CreateChat(ctx, req) -} - -func (c *sdkChatClient) StreamChat(ctx context.Context, chatID uuid.UUID, opts *codersdk.StreamChatOptions) (<-chan codersdk.ChatStreamEvent, io.Closer, error) { - return c.client.StreamChat(ctx, chatID, opts) -} - -func (c *sdkChatClient) CreateChatMessage(ctx context.Context, chatID uuid.UUID, req codersdk.CreateChatMessageRequest) (codersdk.CreateChatMessageResponse, error) { - return c.client.CreateChatMessage(ctx, chatID, req) -} - -func (c *sdkChatClient) UpdateChat(ctx context.Context, chatID uuid.UUID, req codersdk.UpdateChatRequest) error { - return c.client.UpdateChat(ctx, chatID, req) -} - -var _ chatClient = (*sdkChatClient)(nil) +var _ chatModelConfigClient = (*codersdk.ExperimentalClient)(nil) diff --git a/scaletest/chat/config.go b/scaletest/chat/config.go index 5b6b36baa2..703f1c1be6 100644 --- a/scaletest/chat/config.go +++ b/scaletest/chat/config.go @@ -14,6 +14,7 @@ type Config struct { OrganizationID uuid.UUID `json:"organization_id"` // WorkspaceID is the pre-existing workspace to use for this chat run. + // When empty, the chat runs without workspace context. WorkspaceID uuid.UUID `json:"workspace_id"` // Prompt is the text content sent on every turn. @@ -47,9 +48,6 @@ func (c Config) Validate() error { if c.OrganizationID == uuid.Nil { return xerrors.Errorf("validate organization_id: must not be empty") } - if c.WorkspaceID == uuid.Nil { - return xerrors.Errorf("validate workspace_id: must not be empty") - } if c.Prompt == "" { return xerrors.Errorf("validate prompt: must not be empty") } diff --git a/scaletest/chat/provider.go b/scaletest/chat/provider.go index ba946d7db2..156339d8bb 100644 --- a/scaletest/chat/provider.go +++ b/scaletest/chat/provider.go @@ -3,6 +3,7 @@ package chat import ( "context" "net/http" + "time" "github.com/google/uuid" "golang.org/x/xerrors" @@ -12,56 +13,95 @@ import ( ) const ( - scaletestProviderType = "openai-compat" - scaletestProviderDisplayName = "Scaletest LLM Mock" - scaletestModelName = "scaletest-model" - scaletestModelDisplayName = "Scaletest Model" + scaletestAIProviderType = codersdk.AIProviderTypeOpenAICompat + scaletestAIProviderName = "coder-scaletest-mock" + scaletestAIProviderDisplayName = "Scaletest LLM Mock" + scaletestAIProviderAPIKey = "coder-scaletest" + scaletestModelName = "scaletest-model" + scaletestModelDisplayName = "Scaletest Model" + scaletestModelContextLimit = int64(4096) ) -type scaletestProviderAction string +// DefaultProviderPropagationWait is how long to wait after creating or +// updating the mock LLM provider before starting chats. Provider config is +// cached per coderd replica with a 10 second TTL (see +// coderd/x/chatd/configcache.go), and a change is only guaranteed to be +// visible everywhere once every replica's cached entry has expired. 15 +// seconds comfortably exceeds that TTL. +const DefaultProviderPropagationWait = 15 * time.Second + +type scaletestAIProviderAction string const ( - scaletestProviderActionCreated scaletestProviderAction = "created" - scaletestProviderActionUpdated scaletestProviderAction = "updated" - scaletestProviderActionReused scaletestProviderAction = "reused" + scaletestAIProviderActionCreated scaletestAIProviderAction = "created" + scaletestAIProviderActionUpdated scaletestAIProviderAction = "updated" + scaletestAIProviderActionReused scaletestAIProviderAction = "reused" ) -// EnsureScaletestModelConfig bootstraps the shared chat provider and model -// config used by chat scaletests. -func EnsureScaletestModelConfig(ctx context.Context, client *codersdk.ExperimentalClient, logger slog.Logger, llmMockURL string) (uuid.UUID, error) { +// EnsureScaletestModelConfig bootstraps the shared AI provider and model +// config used by chat scaletests. When the provider was created or updated, +// it sleeps for propagationWait so every coderd replica's cached provider +// config expires before chats start. +func EnsureScaletestModelConfig(ctx context.Context, client *codersdk.Client, logger slog.Logger, llmMockURL string, propagationWait time.Duration) (uuid.UUID, error) { + expClient := codersdk.NewExperimentalClient(client) + logger.Info(ctx, "bootstrapping mock LLM provider", slog.F("llm_mock_url", llmMockURL)) - provider, providerAction, err := ensureScaletestProvider(ctx, client, llmMockURL) + provider, providerAction, err := ensureScaletestAIProvider(ctx, expClient, llmMockURL) if err != nil { return uuid.Nil, err } switch providerAction { - case scaletestProviderActionCreated: + case scaletestAIProviderActionCreated: logger.Info(ctx, "created mock LLM provider", - slog.F("provider_type", scaletestProviderType), - slog.F("llm_mock_url", llmMockURL), - ) - case scaletestProviderActionUpdated: - logger.Info(ctx, "updated mock LLM provider", - slog.F("provider_type", scaletestProviderType), + slog.F("provider_name", provider.Name), slog.F("provider_id", provider.ID), slog.F("llm_mock_url", llmMockURL), ) - case scaletestProviderActionReused: + case scaletestAIProviderActionUpdated: + logger.Info(ctx, "updated mock LLM provider", + slog.F("provider_name", provider.Name), + slog.F("provider_id", provider.ID), + slog.F("llm_mock_url", llmMockURL), + ) + case scaletestAIProviderActionReused: logger.Info(ctx, "reusing mock LLM provider", - slog.F("provider_type", scaletestProviderType), + slog.F("provider_name", provider.Name), slog.F("provider_id", provider.ID), ) } + modelConfigID, err := ensureScaletestChatModelConfig(ctx, expClient, logger, provider) + if err != nil { + return uuid.Nil, err + } + + if providerAction != scaletestAIProviderActionReused && propagationWait > 0 { + logger.Info(ctx, "waiting for mock LLM provider propagation", + slog.F("provider_name", provider.Name), + slog.F("wait", propagationWait), + ) + select { + case <-ctx.Done(): + return uuid.Nil, ctx.Err() + case <-time.After(propagationWait): + } + } + + return modelConfigID, nil +} + +func ensureScaletestChatModelConfig(ctx context.Context, client chatModelConfigClient, logger slog.Logger, provider codersdk.AIProvider) (uuid.UUID, error) { modelConfigs, err := client.ListChatModelConfigs(ctx) if err != nil { return uuid.Nil, xerrors.Errorf("list chat model configs: %w", err) } for i := range modelConfigs { - if modelConfigs[i].Provider != provider.Provider || modelConfigs[i].Model != scaletestModelName { + matchesProvider := modelConfigs[i].AIProviderID != nil && *modelConfigs[i].AIProviderID == provider.ID + matchesModel := modelConfigs[i].Model == scaletestModelName + if !matchesProvider || !matchesModel { continue } if !modelConfigs[i].Enabled { @@ -74,9 +114,9 @@ func EnsureScaletestModelConfig(ctx context.Context, client *codersdk.Experiment enabled := true isDefault := false - contextLimit := int64(4096) + contextLimit := scaletestModelContextLimit created, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: provider.Provider, + AIProviderID: &provider.ID, Model: scaletestModelName, DisplayName: scaletestModelDisplayName, Enabled: &enabled, @@ -90,59 +130,66 @@ func EnsureScaletestModelConfig(ctx context.Context, client *codersdk.Experiment return created.ID, nil } -func ensureScaletestProvider(ctx context.Context, client *codersdk.ExperimentalClient, llmMockURL string) (codersdk.ChatProviderConfig, scaletestProviderAction, error) { - enabled := true - mockProviderToken := uuid.NewString() - created, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{ - Provider: scaletestProviderType, - DisplayName: scaletestProviderDisplayName, - APIKey: mockProviderToken, - BaseURL: llmMockURL, - Enabled: &enabled, - }) - if err == nil { - return created, scaletestProviderActionCreated, nil - } - - var sdkErr *codersdk.Error - if !xerrors.As(err, &sdkErr) || sdkErr.StatusCode() != http.StatusConflict { - return codersdk.ChatProviderConfig{}, "", xerrors.Errorf("create scaletest chat provider: %w", err) - } - - providers, err := client.ListChatProviders(ctx) +func ensureScaletestAIProvider(ctx context.Context, client *codersdk.ExperimentalClient, llmMockURL string) (codersdk.AIProvider, scaletestAIProviderAction, error) { + provider, err := client.AIProvider(ctx, scaletestAIProviderName) if err != nil { - return codersdk.ChatProviderConfig{}, "", xerrors.Errorf("list chat providers: %w", err) - } + var sdkErr *codersdk.Error + if !xerrors.As(err, &sdkErr) || sdkErr.StatusCode() != http.StatusNotFound { + return codersdk.AIProvider{}, "", xerrors.Errorf("look up scaletest AI provider: %w", err) + } - var existing *codersdk.ChatProviderConfig - for i := range providers { - if providers[i].Provider == scaletestProviderType { - existing = &providers[i] - break + created, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{ + Type: scaletestAIProviderType, + Name: scaletestAIProviderName, + DisplayName: scaletestAIProviderDisplayName, + Enabled: true, + BaseURL: llmMockURL, + APIKeys: []string{scaletestAIProviderAPIKey}, + }) + if err == nil { + return created, scaletestAIProviderActionCreated, nil + } + + sdkErr = nil + if !xerrors.As(err, &sdkErr) || sdkErr.StatusCode() != http.StatusConflict { + return codersdk.AIProvider{}, "", xerrors.Errorf("create scaletest AI provider: %w", err) + } + + provider, err = client.AIProvider(ctx, scaletestAIProviderName) + if err != nil { + return codersdk.AIProvider{}, "", xerrors.Errorf("look up scaletest AI provider after conflict: %w", err) } } - if existing == nil { - return codersdk.ChatProviderConfig{}, "", xerrors.Errorf("find existing %s provider after conflict: not found", scaletestProviderType) + + if provider.Type != scaletestAIProviderType { + return codersdk.AIProvider{}, "", xerrors.Errorf("refusing to use scaletest AI provider %s with type %q", provider.ID, provider.Type) } - if existing.DisplayName != scaletestProviderDisplayName { - return codersdk.ChatProviderConfig{}, "", xerrors.Errorf("refusing to overwrite existing %s provider %s with display name %q", scaletestProviderType, existing.ID, existing.DisplayName) + if provider.DisplayName != scaletestAIProviderDisplayName { + return codersdk.AIProvider{}, "", xerrors.Errorf("refusing to use scaletest AI provider %s with display name %q", provider.ID, provider.DisplayName) + } + if !provider.Enabled { + return codersdk.AIProvider{}, "", xerrors.Errorf("existing scaletest AI provider %s is disabled; re-enable or delete it before running scaletests", provider.ID) } - if !existing.Enabled { - return codersdk.ChatProviderConfig{}, "", xerrors.Errorf("existing scaletest chat provider %s is disabled; re-enable or delete it before running scaletests", existing.ID) + var update codersdk.UpdateAIProviderRequest + needsUpdate := false + if provider.BaseURL != llmMockURL { + update.BaseURL = &llmMockURL + needsUpdate = true } - if existing.BaseURL == llmMockURL { - return *existing, scaletestProviderActionReused, nil + if len(provider.APIKeys) == 0 { + apiKey := scaletestAIProviderAPIKey + apiKeys := []codersdk.AIProviderKeyMutation{{APIKey: &apiKey}} + update.APIKeys = &apiKeys + needsUpdate = true + } + if !needsUpdate { + return provider, scaletestAIProviderActionReused, nil } - updated, err := client.UpdateChatProvider(ctx, existing.ID, codersdk.UpdateChatProviderConfigRequest{ - DisplayName: scaletestProviderDisplayName, - APIKey: &mockProviderToken, - BaseURL: &llmMockURL, - Enabled: &enabled, - }) + updated, err := client.UpdateAIProvider(ctx, scaletestAIProviderName, update) if err != nil { - return codersdk.ChatProviderConfig{}, "", xerrors.Errorf("update scaletest chat provider: %w", err) + return codersdk.AIProvider{}, "", xerrors.Errorf("update scaletest AI provider: %w", err) } - return updated, scaletestProviderActionUpdated, nil + return updated, scaletestAIProviderActionUpdated, nil } diff --git a/scaletest/chat/run.go b/scaletest/chat/run.go index b2e591fab6..d5b98d6381 100644 --- a/scaletest/chat/run.go +++ b/scaletest/chat/run.go @@ -54,7 +54,7 @@ var ( func NewRunner(client *codersdk.Client, cfg Config) *Runner { return &Runner{ - client: newChatClient(client), + client: codersdk.NewExperimentalClient(client), cfg: cfg, } } @@ -108,15 +108,18 @@ func (r *Runner) Run(ctx context.Context, id string, logs io.Writer) error { r.resetConversation(time.Now(), markTurnStartReady) createStartedAt := time.Now() - chat, err := r.client.CreateChat(ctx, codersdk.CreateChatRequest{ + createReq := codersdk.CreateChatRequest{ OrganizationID: r.cfg.OrganizationID, - WorkspaceID: &workspaceID, ModelConfigID: &modelConfigID, Content: []codersdk.ChatInputPart{{ Type: codersdk.ChatInputPartTypeText, Text: r.cfg.Prompt, }}, - }) + } + if workspaceID != uuid.Nil { + createReq.WorkspaceID = &workspaceID + } + chat, err := r.client.CreateChat(ctx, createReq) if err != nil { r.result.failureStage = failureStageCreateChat r.cfg.Metrics.ChatStageFailuresTotal.WithLabelValues(r.result.failureStage).Inc()