Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion cmd/kael/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,7 @@ func main() {

func run(settings config.Config, logger *zap.Logger) error {
tlsVerify := !settings.IgnoreVerifyCerts
componentClient, err := component.Connect(component.Options{CoreURL: settings.CoreHost, TLSVerify: tlsVerify, Timeout: settings.HTTPRequestTimeout, Name: settings.Name, BootstrapToken: settings.BootstrapToken, AccessKeyFile: settings.AccessKeyFilePath})
componentClient, err := component.Connect(component.Options{CoreURL: settings.CoreHost, TLSVerify: tlsVerify, Timeout: settings.HTTPRequestTimeout, Name: settings.Name, BootstrapToken: settings.BootstrapToken, AccessKeyFile: settings.AccessKeyFilePath, Logger: logger})
if err != nil {
return err
}
Expand Down
44 changes: 40 additions & 4 deletions internal/component/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,8 @@ const (
runtimeStorePath = "/api/v1/chat-ai/runtime-store/"
openAPIPath = "/api/swagger.json"
maxOpenAPIBytes = int64(32 * 1024 * 1024)
connectAttempts = 10
connectRetryDelay = 5 * time.Second
)

var errInvalidAccessKey = errors.New("component access key file is invalid")
Expand All @@ -52,6 +54,7 @@ type Options struct {
Name string
BootstrapToken string
AccessKeyFile string
Logger *zap.Logger
}

type Client struct {
Expand Down Expand Up @@ -115,9 +118,16 @@ type runtimeStoreAppendResponse struct {
}

func Connect(options Options) (*Client, error) {
return connect(options, time.Sleep)
}

func connect(options Options, wait func(time.Duration)) (*Client, error) {
if strings.TrimSpace(options.Name) == "" || strings.TrimSpace(options.AccessKeyFile) == "" {
return nil, fmt.Errorf("component identity is incomplete")
}
if options.Logger == nil {
options.Logger = zap.NewNop()
}
if options.Timeout < 30*time.Second {
options.Timeout = 30 * time.Second
}
Expand All @@ -133,16 +143,28 @@ func Connect(options Options) (*Client, error) {
if clientErr != nil {
return nil, clientErr
}
if valid, validationErr := validateAccessKey(client); validationErr == nil && valid {
return connectedClient(options, client, key), nil
} else if validationErr != nil {
var valid bool
validationErr := retryConnect(options.Logger, wait, func() error {
var err error
valid, err = validateAccessKey(client)
return err
})
if validationErr != nil {
return nil, validationErr
}
if valid {
return connectedClient(options, client, key), nil
}
options.Logger.Warn("Kael component access key unauthorized; registering a new access key")
}
if strings.TrimSpace(options.BootstrapToken) == "" {
return nil, fmt.Errorf("component access key is missing or unauthorized; set BOOTSTRAP_TOKEN to register Kael")
}
key, err = register(options)
err = retryConnect(options.Logger, wait, func() error {
var registerErr error
key, registerErr = register(options)
return registerErr
})
if err != nil {
return nil, err
}
Expand All @@ -156,6 +178,20 @@ func Connect(options Options) (*Client, error) {
return connectedClient(options, client, key), nil
}

func retryConnect(logger *zap.Logger, wait func(time.Duration), request func() error) error {
var err error
for attempt := 1; attempt <= connectAttempts; attempt++ {
if err = request(); err == nil {
return nil
}
logger.Warn("Core component request failed", zap.Int("attempt", attempt), zap.Int("max_attempts", connectAttempts), zap.Error(err))
if attempt < connectAttempts {
wait(connectRetryDelay)
}
}
return fmt.Errorf("connect to Core failed after %d attempts: %w", connectAttempts, err)
}

func connectedClient(options Options, client *httplib.Client, key accessKey) *Client {
return &Client{
client: client,
Expand Down
64 changes: 56 additions & 8 deletions internal/component/client_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -71,26 +71,74 @@ func TestComponentRegistrationModelConfigAndHeartbeat(t *testing.T) {
}
}

func TestComponentReusesValidAccessKey(t *testing.T) {
func TestComponentConnectRetryLimit(t *testing.T) {
for _, path := range []string{registerPath, profilePath} {
t.Run(path, func(t *testing.T) {
attempts := 0
server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
if request.URL.Path != path {
t.Errorf("unexpected request path: %s", request.URL.Path)
}
attempts++
response.WriteHeader(http.StatusServiceUnavailable)
}))
defer server.Close()

keyPath := filepath.Join(t.TempDir(), ".access_key")
if path == profilePath {
if err := os.WriteFile(keyPath, []byte("access-id:access-secret"), 0o600); err != nil {
t.Fatal(err)
}
}
options := Options{CoreURL: server.URL, Name: "kael-test", BootstrapToken: "bootstrap-secret", AccessKeyFile: keyPath}
_, err := connect(options, func(time.Duration) {})
if err == nil || attempts != 10 {
t.Fatalf("expected failure after 10 attempts: attempts=%d err=%v", attempts, err)
}
})
}
}

func TestComponentReregistersUnauthorizedAccessKey(t *testing.T) {
registrations := 0
server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
response.Header().Set("Content-Type", "application/json")
if request.URL.Path != profilePath {
t.Errorf("unexpected request path: %s", request.URL.Path)
switch request.URL.Path {
case profilePath:
assertSigned(t, request)
if !strings.Contains(request.Header.Get("Authorization"), `keyId="access-id"`) {
response.WriteHeader(http.StatusUnauthorized)
return
}
_, _ = response.Write([]byte(`{"id":"service-account-id"}`))
case registerPath:
registrations++
if registrations == 1 {
response.WriteHeader(http.StatusServiceUnavailable)
return
}
_, _ = response.Write([]byte(`{"service_account":{"access_key":{"id":"access-id","secret":"access-secret"}}}`))
default:
http.NotFound(response, request)
return
}
assertSigned(t, request)
_, _ = response.Write([]byte(`{"id":"service-account-id"}`))
}))
defer server.Close()

keyPath := filepath.Join(t.TempDir(), ".access_key")
if err := os.WriteFile(keyPath, []byte("access-id:access-secret"), 0o600); err != nil {
if err := os.WriteFile(keyPath, []byte("expired-id:expired-secret"), 0o600); err != nil {
t.Fatal(err)
}
if _, err := Connect(Options{CoreURL: server.URL, TLSVerify: true, Name: "kael-test", AccessKeyFile: keyPath}); err != nil {
options := Options{CoreURL: server.URL, Name: "kael-test", BootstrapToken: "bootstrap-secret", AccessKeyFile: keyPath}
if _, err := connect(options, func(time.Duration) {}); err != nil {
t.Fatal(err)
}
options.BootstrapToken = ""
if _, err := Connect(options); err != nil {
t.Fatalf("restart did not reuse the new access key: %v", err)
}
if registrations != 2 {
t.Fatalf("expected registration to succeed on retry: registrations=%d", registrations)
}
}

func assertSigned(t *testing.T, request *http.Request) {
Expand Down
Loading