mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add experimental agents support (#22290)
feat: add AI chat system with agent tools and chat UI Introduce the chatd subsystem and Agents UI for AI-powered chat within Coder workspaces. - Add chatd package with chat loop, message compaction, prompt management, and LLM provider integration (OpenAI, Anthropic) - Add agent tools: create workspace, list/read templates, read/write/ edit files, execute commands - Add chat API endpoints with streaming, message editing, and durable reconnection - Add database schema and migrations for chats, chat messages, chat providers, and chat model configs - Add RBAC policies and dbauthz enforcement for chat resources - Add Agents UI pages with conversation timeline, queued messages list, diff viewer, and model configuration panel - Add comprehensive test coverage including coderd integration tests, chatd unit tests, and Storybook stories - Gate feature behind experiments flag --------- Co-authored-by: Cian Johnston <cian@coder.com> Co-authored-by: Danielle Maywood <danielle@themaywoods.com> Co-authored-by: Jeremy Ruppel <jeremy@coder.com> Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Cian Johnston
Danielle Maywood
Jeremy Ruppel
Claude Sonnet 4.6
parent
67da4e8b56
commit
edee917d88
@@ -0,0 +1,172 @@
|
||||
package coderd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/url"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/chatd"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/websocket"
|
||||
)
|
||||
|
||||
// RelaySourceHeader marks replica-relayed stream requests.
|
||||
const RelaySourceHeader = "X-Coder-Relay-Source-Replica"
|
||||
|
||||
const (
|
||||
authorizationHeader = "Authorization"
|
||||
cookieHeader = "Cookie"
|
||||
)
|
||||
|
||||
// newRemotePartsProvider creates a RemotePartsProvider that dials a remote
|
||||
// replica's stream endpoint to fetch message_part events. It filters to only
|
||||
// forward message_part events since durable events come via pubsub.
|
||||
func newRemotePartsProvider(
|
||||
resolveReplicaAddress func(context.Context, uuid.UUID) (string, bool),
|
||||
replicaHTTPClient *http.Client,
|
||||
replicaID uuid.UUID,
|
||||
) chatd.RemotePartsProvider {
|
||||
return func(
|
||||
ctx context.Context,
|
||||
chatID uuid.UUID,
|
||||
workerID uuid.UUID,
|
||||
requestHeader http.Header,
|
||||
) (
|
||||
[]codersdk.ChatStreamEvent,
|
||||
<-chan codersdk.ChatStreamEvent,
|
||||
func(),
|
||||
error,
|
||||
) {
|
||||
address, ok := resolveReplicaAddress(ctx, workerID)
|
||||
if !ok {
|
||||
return nil, nil, nil, xerrors.New("worker replica not found")
|
||||
}
|
||||
|
||||
baseURL, err := url.Parse(address)
|
||||
if err != nil {
|
||||
return nil, nil, nil, xerrors.Errorf("parse relay address %q: %w", address, err)
|
||||
}
|
||||
relayCtx, relayCancel := context.WithCancel(ctx)
|
||||
sdkClient := codersdk.New(baseURL)
|
||||
sdkClient.HTTPClient = replicaHTTPClient
|
||||
sdkClient.SessionTokenProvider = relayHeaderTokenProvider{
|
||||
header: relayHeaders(requestHeader, replicaID),
|
||||
}
|
||||
sourceEvents, sourceStream, err := sdkClient.StreamChat(relayCtx, chatID)
|
||||
if err != nil {
|
||||
relayCancel()
|
||||
return nil, nil, nil, xerrors.Errorf("dial relay stream: %w", err)
|
||||
}
|
||||
|
||||
snapshot := make([]codersdk.ChatStreamEvent, 0, 100)
|
||||
preloaded := make([]codersdk.ChatStreamEvent, 0, 100)
|
||||
drainInitial:
|
||||
for len(snapshot) < cap(snapshot) {
|
||||
select {
|
||||
case <-relayCtx.Done():
|
||||
_ = sourceStream.Close()
|
||||
relayCancel()
|
||||
return nil, nil, nil, xerrors.Errorf("dial relay stream: %w", relayCtx.Err())
|
||||
case event, ok := <-sourceEvents:
|
||||
if !ok {
|
||||
break drainInitial
|
||||
}
|
||||
if event.Type != codersdk.ChatStreamEventTypeMessagePart {
|
||||
continue
|
||||
}
|
||||
snapshot = append(snapshot, event)
|
||||
preloaded = append(preloaded, event)
|
||||
default:
|
||||
break drainInitial
|
||||
}
|
||||
}
|
||||
|
||||
events := make(chan codersdk.ChatStreamEvent, 128)
|
||||
|
||||
go func() {
|
||||
defer close(events)
|
||||
defer relayCancel()
|
||||
defer func() {
|
||||
_ = sourceStream.Close()
|
||||
}()
|
||||
|
||||
for _, event := range preloaded {
|
||||
select {
|
||||
case events <- event:
|
||||
case <-relayCtx.Done():
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-relayCtx.Done():
|
||||
return
|
||||
case event, ok := <-sourceEvents:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if event.Type != codersdk.ChatStreamEventTypeMessagePart {
|
||||
continue
|
||||
}
|
||||
select {
|
||||
case events <- event:
|
||||
case <-relayCtx.Done():
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
cancel := func() {
|
||||
relayCancel()
|
||||
_ = sourceStream.Close()
|
||||
}
|
||||
return snapshot, events, cancel, nil
|
||||
}
|
||||
}
|
||||
|
||||
type relayHeaderTokenProvider struct {
|
||||
header http.Header
|
||||
}
|
||||
|
||||
func (p relayHeaderTokenProvider) AsRequestOption() codersdk.RequestOption {
|
||||
return func(req *http.Request) {
|
||||
for key, values := range p.header {
|
||||
for _, value := range values {
|
||||
req.Header.Add(key, value)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p relayHeaderTokenProvider) SetDialOption(opts *websocket.DialOptions) {
|
||||
if opts.HTTPHeader == nil {
|
||||
opts.HTTPHeader = make(http.Header)
|
||||
}
|
||||
for key, values := range p.header {
|
||||
for _, value := range values {
|
||||
opts.HTTPHeader.Add(key, value)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p relayHeaderTokenProvider) GetSessionToken() string {
|
||||
return p.header.Get(codersdk.SessionTokenHeader)
|
||||
}
|
||||
|
||||
func relayHeaders(source http.Header, replicaID uuid.UUID) http.Header {
|
||||
header := make(http.Header)
|
||||
if source != nil {
|
||||
for _, key := range []string{codersdk.SessionTokenHeader, authorizationHeader, cookieHeader} {
|
||||
for _, value := range source.Values(key) {
|
||||
header.Add(key, value)
|
||||
}
|
||||
}
|
||||
}
|
||||
header.Set(RelaySourceHeader, replicaID.String())
|
||||
return header
|
||||
}
|
||||
@@ -0,0 +1,355 @@
|
||||
package coderd_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/url"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/chatd/chattest"
|
||||
"github.com/coder/coder/v2/coderd/coderdtest"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/enterprise/coderd/coderdenttest"
|
||||
"github.com/coder/coder/v2/enterprise/coderd/license"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestChatStreamRelay(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("RelayMessagePartsAcrossReplicas", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
db, pubsub := dbtestutil.NewDB(t)
|
||||
firstClient, _ := coderdenttest.New(t, &coderdenttest.Options{
|
||||
Options: &coderdtest.Options{
|
||||
Database: db,
|
||||
Pubsub: pubsub,
|
||||
},
|
||||
LicenseOptions: &coderdenttest.LicenseOptions{
|
||||
Features: license.Features{
|
||||
codersdk.FeatureHighAvailability: 1,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
secondClient, _ := coderdenttest.New(t, &coderdenttest.Options{
|
||||
Options: &coderdtest.Options{
|
||||
Database: db,
|
||||
Pubsub: pubsub,
|
||||
},
|
||||
DontAddLicense: true,
|
||||
DontAddFirstUser: true,
|
||||
})
|
||||
secondClient.SetSessionToken(firstClient.SessionToken())
|
||||
|
||||
// Verify we have two replicas
|
||||
replicas, err := secondClient.Replicas(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, replicas, 2)
|
||||
firstReplicaID := replicaIDForClientURL(t, firstClient.URL, replicas)
|
||||
secondReplicaID := replicaIDForClientURL(t, secondClient.URL, replicas)
|
||||
|
||||
streamingChunks := make(chan chattest.OpenAIChunk, 8)
|
||||
chatStreamStarted := make(chan struct{}, 1)
|
||||
openai := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if req.Stream {
|
||||
select {
|
||||
case chatStreamStarted <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
return chattest.OpenAIResponse{StreamingChunks: streamingChunks}
|
||||
}
|
||||
return chattest.OpenAINonStreamingResponse("ok")
|
||||
})
|
||||
|
||||
//nolint:gocritic // Test uses owner client to configure chat providers.
|
||||
provider, err := firstClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
|
||||
Provider: "openai",
|
||||
DisplayName: "OpenAI",
|
||||
APIKey: "test",
|
||||
BaseURL: openai,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, codersdk.ChatProviderConfigSourceDatabase, provider.Source)
|
||||
|
||||
model, err := firstClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: provider.Provider,
|
||||
Model: "gpt-4",
|
||||
DisplayName: "GPT-4",
|
||||
ContextLimit: &[]int64{1000}[0],
|
||||
CompressionThreshold: &[]int32{70}[0],
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create a chat on the first replica
|
||||
chat, err := firstClient.CreateChat(ctx, codersdk.CreateChatRequest{
|
||||
Content: []codersdk.ChatInputPart{{
|
||||
Type: codersdk.ChatInputPartTypeText,
|
||||
Text: "Test chat for relay",
|
||||
}},
|
||||
ModelConfigID: &model.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, codersdk.ChatStatusPending, chat.Status)
|
||||
|
||||
var runningChat database.Chat
|
||||
require.Eventually(t, func() bool {
|
||||
current, getErr := db.GetChatByID(ctx, chat.ID)
|
||||
if getErr != nil {
|
||||
return false
|
||||
}
|
||||
if current.Status != database.ChatStatusRunning || !current.WorkerID.Valid {
|
||||
return false
|
||||
}
|
||||
runningChat = current
|
||||
return true
|
||||
}, testutil.WaitLong, testutil.IntervalFast)
|
||||
|
||||
var localClient *codersdk.Client
|
||||
var relayClient *codersdk.Client
|
||||
switch runningChat.WorkerID.UUID {
|
||||
case firstReplicaID:
|
||||
localClient = firstClient
|
||||
relayClient = secondClient
|
||||
case secondReplicaID:
|
||||
localClient = secondClient
|
||||
relayClient = firstClient
|
||||
default:
|
||||
require.FailNowf(
|
||||
t,
|
||||
"worker replica was not recognized",
|
||||
"worker %s was not one of %s or %s",
|
||||
runningChat.WorkerID.UUID,
|
||||
firstReplicaID,
|
||||
secondReplicaID,
|
||||
)
|
||||
}
|
||||
|
||||
firstEvents, firstStream, err := localClient.StreamChat(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
defer firstStream.Close()
|
||||
|
||||
select {
|
||||
case <-chatStreamStarted:
|
||||
case <-ctx.Done():
|
||||
require.FailNowf(
|
||||
t,
|
||||
"timed out waiting for OpenAI stream request",
|
||||
"chat stream request did not start before context deadline: %v",
|
||||
ctx.Err(),
|
||||
)
|
||||
}
|
||||
|
||||
firstChunkText := "relay-part-one"
|
||||
streamingChunks <- chattest.OpenAITextChunks(firstChunkText)[0]
|
||||
firstEvent := waitForStreamTextPart(ctx, t, firstEvents, firstChunkText)
|
||||
require.Equal(t, "assistant", firstEvent.MessagePart.Role)
|
||||
|
||||
secondEvents, secondStream, err := relayClient.StreamChat(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
defer secondStream.Close()
|
||||
|
||||
secondSnapshotEvent := waitForStreamTextPart(ctx, t, secondEvents, firstChunkText)
|
||||
require.Equal(t, "assistant", secondSnapshotEvent.MessagePart.Role)
|
||||
|
||||
secondChunkText := "relay-part-two"
|
||||
streamingChunks <- chattest.OpenAITextChunks(secondChunkText)[0]
|
||||
waitForStreamTextPart(ctx, t, firstEvents, secondChunkText)
|
||||
waitForStreamTextPart(ctx, t, secondEvents, secondChunkText)
|
||||
|
||||
close(streamingChunks)
|
||||
})
|
||||
}
|
||||
|
||||
func waitForStreamTextPart(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
events <-chan codersdk.ChatStreamEvent,
|
||||
expectedText string,
|
||||
) codersdk.ChatStreamEvent {
|
||||
t.Helper()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
require.FailNowf(
|
||||
t,
|
||||
"timed out waiting for chat stream event",
|
||||
"expected text part %q before context deadline: %v",
|
||||
expectedText,
|
||||
ctx.Err(),
|
||||
)
|
||||
case event, ok := <-events:
|
||||
require.Truef(t, ok, "chat stream closed while waiting for %q", expectedText)
|
||||
|
||||
if event.Type == codersdk.ChatStreamEventTypeError {
|
||||
errMessage := "unknown chat stream error"
|
||||
if event.Error != nil && event.Error.Message != "" {
|
||||
errMessage = event.Error.Message
|
||||
}
|
||||
require.FailNowf(
|
||||
t,
|
||||
"chat stream returned error event",
|
||||
"while waiting for %q: %s",
|
||||
expectedText,
|
||||
errMessage,
|
||||
)
|
||||
}
|
||||
|
||||
if event.Type != codersdk.ChatStreamEventTypeMessagePart || event.MessagePart == nil {
|
||||
continue
|
||||
}
|
||||
if event.MessagePart.Part.Type != codersdk.ChatMessagePartTypeText {
|
||||
continue
|
||||
}
|
||||
|
||||
require.Equal(t, expectedText, event.MessagePart.Part.Text)
|
||||
return event
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func replicaIDForClientURL(
|
||||
t *testing.T,
|
||||
clientURL *url.URL,
|
||||
replicas []codersdk.Replica,
|
||||
) uuid.UUID {
|
||||
t.Helper()
|
||||
|
||||
for _, replica := range replicas {
|
||||
relayURL, err := url.Parse(replica.RelayAddress)
|
||||
require.NoErrorf(
|
||||
t,
|
||||
err,
|
||||
"parse replica relay address %q",
|
||||
replica.RelayAddress,
|
||||
)
|
||||
if relayURL.Host == clientURL.Host {
|
||||
return replica.ID
|
||||
}
|
||||
}
|
||||
|
||||
require.FailNowf(
|
||||
t,
|
||||
"missing replica for client URL",
|
||||
"client host %q not present in replica list",
|
||||
clientURL.Host,
|
||||
)
|
||||
return uuid.Nil
|
||||
}
|
||||
|
||||
func TestChatModelConfigDefault(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
client, _ := coderdenttest.New(t, nil)
|
||||
|
||||
//nolint:gocritic // Test uses owner client to configure chat providers.
|
||||
provider, err := client.CreateChatProvider(
|
||||
ctx,
|
||||
codersdk.CreateChatProviderConfigRequest{
|
||||
Provider: "openai",
|
||||
DisplayName: "OpenAI",
|
||||
APIKey: "test",
|
||||
BaseURL: "https://example.com",
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
contextLimit := int64(1000)
|
||||
compressionThreshold := int32(70)
|
||||
trueValue := true
|
||||
falseValue := false
|
||||
|
||||
firstModel, err := client.CreateChatModelConfig(
|
||||
ctx,
|
||||
codersdk.CreateChatModelConfigRequest{
|
||||
Provider: provider.Provider,
|
||||
Model: "gpt-5-a",
|
||||
DisplayName: "GPT 5 A",
|
||||
IsDefault: &trueValue,
|
||||
ContextLimit: &contextLimit,
|
||||
CompressionThreshold: &compressionThreshold,
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, firstModel.IsDefault)
|
||||
|
||||
secondModel, err := client.CreateChatModelConfig(
|
||||
ctx,
|
||||
codersdk.CreateChatModelConfigRequest{
|
||||
Provider: provider.Provider,
|
||||
Model: "gpt-5-b",
|
||||
DisplayName: "GPT 5 B",
|
||||
IsDefault: &trueValue,
|
||||
ContextLimit: &contextLimit,
|
||||
CompressionThreshold: &compressionThreshold,
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, secondModel.IsDefault)
|
||||
|
||||
modelConfigs, err := client.ListChatModelConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
firstStored := findChatModelConfigByID(t, modelConfigs, firstModel.ID)
|
||||
secondStored := findChatModelConfigByID(t, modelConfigs, secondModel.ID)
|
||||
require.False(t, firstStored.IsDefault)
|
||||
require.True(t, secondStored.IsDefault)
|
||||
|
||||
updatedFirst, err := client.UpdateChatModelConfig(
|
||||
ctx,
|
||||
firstModel.ID,
|
||||
codersdk.UpdateChatModelConfigRequest{
|
||||
IsDefault: &trueValue,
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.True(t, updatedFirst.IsDefault)
|
||||
|
||||
modelConfigs, err = client.ListChatModelConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
firstStored = findChatModelConfigByID(t, modelConfigs, firstModel.ID)
|
||||
secondStored = findChatModelConfigByID(t, modelConfigs, secondModel.ID)
|
||||
require.True(t, firstStored.IsDefault)
|
||||
require.False(t, secondStored.IsDefault)
|
||||
|
||||
updatedFirst, err = client.UpdateChatModelConfig(
|
||||
ctx,
|
||||
firstModel.ID,
|
||||
codersdk.UpdateChatModelConfigRequest{
|
||||
IsDefault: &falseValue,
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.False(t, updatedFirst.IsDefault)
|
||||
|
||||
modelConfigs, err = client.ListChatModelConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
firstStored = findChatModelConfigByID(t, modelConfigs, firstModel.ID)
|
||||
secondStored = findChatModelConfigByID(t, modelConfigs, secondModel.ID)
|
||||
require.False(t, firstStored.IsDefault)
|
||||
require.True(t, secondStored.IsDefault)
|
||||
}
|
||||
|
||||
func findChatModelConfigByID(
|
||||
t *testing.T,
|
||||
modelConfigs []codersdk.ChatModelConfig,
|
||||
id uuid.UUID,
|
||||
) codersdk.ChatModelConfig {
|
||||
t.Helper()
|
||||
|
||||
for _, modelConfig := range modelConfigs {
|
||||
if modelConfig.ID == id {
|
||||
return modelConfig
|
||||
}
|
||||
}
|
||||
|
||||
require.FailNowf(t, "missing model config", "model config %s not found", id)
|
||||
return codersdk.ChatModelConfig{}
|
||||
}
|
||||
@@ -3,6 +3,7 @@ package coderd
|
||||
import (
|
||||
"context"
|
||||
"crypto/ed25519"
|
||||
"crypto/tls"
|
||||
"fmt"
|
||||
"math"
|
||||
"net/http"
|
||||
@@ -15,6 +16,7 @@ import (
|
||||
|
||||
"github.com/cenkalti/backoff/v4"
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/google/uuid"
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"golang.org/x/xerrors"
|
||||
"tailscale.com/tailcfg"
|
||||
@@ -100,6 +102,11 @@ func New(ctx context.Context, options *Options) (_ *API, err error) {
|
||||
}
|
||||
|
||||
ctx, cancelFunc := context.WithCancel(ctx)
|
||||
defer func() {
|
||||
if err != nil {
|
||||
cancelFunc()
|
||||
}
|
||||
}()
|
||||
|
||||
if options.ExternalTokenEncryption == nil {
|
||||
options.ExternalTokenEncryption = make([]dbcrypt.Cipher, 0)
|
||||
@@ -141,6 +148,33 @@ func New(ctx context.Context, options *Options) (_ *API, err error) {
|
||||
)
|
||||
}
|
||||
|
||||
meshTLSConfig, err := replicasync.CreateDERPMeshTLSConfig(options.AccessURL.Hostname(), options.TLSCertificates)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("create DERP mesh TLS config: %w", err)
|
||||
}
|
||||
|
||||
var replicaManagerPtr atomic.Pointer[replicasync.Manager]
|
||||
resolveReplicaAddress := func(
|
||||
_ context.Context,
|
||||
replicaID uuid.UUID,
|
||||
) (string, bool) {
|
||||
manager := replicaManagerPtr.Load()
|
||||
if manager == nil {
|
||||
return "", false
|
||||
}
|
||||
for _, replica := range manager.AllPrimary() {
|
||||
if replica.ID != replicaID {
|
||||
continue
|
||||
}
|
||||
relayAddress := strings.TrimSpace(replica.RelayAddress)
|
||||
if relayAddress == "" {
|
||||
return "", false
|
||||
}
|
||||
return relayAddress, true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
api := &API{
|
||||
ctx: ctx,
|
||||
cancel: cancelFunc,
|
||||
@@ -156,6 +190,44 @@ func New(ctx context.Context, options *Options) (_ *API, err error) {
|
||||
}
|
||||
// This must happen before coderd initialization!
|
||||
options.PostAuthAdditionalHeadersFunc = api.writeEntitlementWarningsHeader
|
||||
|
||||
// Wire up enterprise chat relay for cross-replica message_part streaming.
|
||||
// Must be set before coderd.New so the chat processor gets it.
|
||||
replicaHTTPClient := replicaRelayHTTPClient(options.HTTPClient, meshTLSConfig)
|
||||
if replicaHTTPClient == nil {
|
||||
replicaHTTPClient = options.Options.HTTPClient
|
||||
}
|
||||
if replicaHTTPClient == nil {
|
||||
replicaHTTPClient = http.DefaultClient
|
||||
}
|
||||
// Use a closure that captures api by reference so it can access api.AGPL.ID
|
||||
// after coderd.New is called. The provider is only invoked when Subscribe
|
||||
// is called, which happens after initialization, so api.AGPL will be set.
|
||||
options.Options.ChatRemotePartsProvider = func(
|
||||
ctx context.Context,
|
||||
chatID uuid.UUID,
|
||||
workerID uuid.UUID,
|
||||
requestHeader http.Header,
|
||||
) (
|
||||
[]codersdk.ChatStreamEvent,
|
||||
<-chan codersdk.ChatStreamEvent,
|
||||
func(),
|
||||
error,
|
||||
) {
|
||||
// Get the replica ID from the API (will be set after coderd.New)
|
||||
replicaID := api.AGPL.ID
|
||||
if replicaID == uuid.Nil {
|
||||
// Fallback if somehow called before initialization
|
||||
replicaID = uuid.New()
|
||||
}
|
||||
provider := newRemotePartsProvider(
|
||||
resolveReplicaAddress,
|
||||
replicaHTTPClient,
|
||||
replicaID,
|
||||
)
|
||||
return provider(ctx, chatID, workerID, requestHeader)
|
||||
}
|
||||
|
||||
api.AGPL = coderd.New(options.Options)
|
||||
defer func() {
|
||||
if err != nil {
|
||||
@@ -583,10 +655,6 @@ func New(ctx context.Context, options *Options) (_ *API, err error) {
|
||||
})))
|
||||
}
|
||||
|
||||
meshTLSConfig, err := replicasync.CreateDERPMeshTLSConfig(options.AccessURL.Hostname(), options.TLSCertificates)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("create DERP mesh TLS config: %w", err)
|
||||
}
|
||||
// We always want to run the replica manager even if we don't have DERP
|
||||
// enabled, since it's used to detect other coder servers for licensing.
|
||||
api.replicaManager, err = replicasync.New(ctx, options.Logger, options.Database, options.Pubsub, &replicasync.Options{
|
||||
@@ -600,6 +668,7 @@ func New(ctx context.Context, options *Options) (_ *API, err error) {
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("initialize replica: %w", err)
|
||||
}
|
||||
replicaManagerPtr.Store(api.replicaManager)
|
||||
if api.DERPServer != nil {
|
||||
api.derpMesh = derpmesh.New(options.Logger.Named("derpmesh"), api.DERPServer, meshTLSConfig)
|
||||
}
|
||||
@@ -651,6 +720,28 @@ func New(ctx context.Context, options *Options) (_ *API, err error) {
|
||||
return api, nil
|
||||
}
|
||||
|
||||
func replicaRelayHTTPClient(base *http.Client, tlsConfig *tls.Config) *http.Client {
|
||||
if base == nil {
|
||||
base = http.DefaultClient
|
||||
}
|
||||
|
||||
clone := *base
|
||||
var transport *http.Transport
|
||||
switch t := base.Transport.(type) {
|
||||
case *http.Transport:
|
||||
transport = t.Clone()
|
||||
default:
|
||||
if defaultTransport, ok := http.DefaultTransport.(*http.Transport); ok {
|
||||
transport = defaultTransport.Clone()
|
||||
} else {
|
||||
transport = &http.Transport{}
|
||||
}
|
||||
}
|
||||
transport.TLSClientConfig = tlsConfig
|
||||
clone.Transport = transport
|
||||
return &clone
|
||||
}
|
||||
|
||||
type Options struct {
|
||||
*coderd.Options
|
||||
|
||||
|
||||
Reference in New Issue
Block a user