From bfaa3340e8a71a326baafaa645dd231ec0857395 Mon Sep 17 00:00:00 2001 From: Jakub Nyckowski Date: Mon, 10 Jun 2024 16:57:05 -0400 Subject: [PATCH] Remove Assist (#42657) * Remove assist feature * Fix some tests * Remove unused functions * Fix ut Remove more stuff --- .github/ISSUE_TEMPLATE/testplan.md | 17 - CHANGELOG.md | 5 +- api/client/client.go | 81 - api/client/webclient/webclient.go | 2 - api/client/webclient/webconfig.go | 2 - api/defaults/defaults.go | 4 - api/gen/proto/go/assist/v1/assist.pb.go | 1595 ----------------- api/gen/proto/go/assist/v1/assist_grpc.pb.go | 509 ------ .../userpreferences/v1/userpreferences.pb.go | 221 ++- api/proto/teleport/assist/v1/assist.proto | 201 --- .../userpreferences/v1/userpreferences.proto | 4 +- api/types/constants.go | 4 - api/types/networking.go | 22 - constants.go | 7 +- .../userpreferences/v1/userpreferences_pb.ts | 14 - go.mod | 1 - go.sum | 2 - .../terraform/tfschema/types_terraform.go | 44 - lib/ai/chat.go | 92 - lib/ai/chat_test.go | 364 ---- lib/ai/client.go | 253 --- lib/ai/client_test.go | 112 -- lib/ai/embedding/embedding.go | 115 -- lib/ai/embedding/serialization.go | 152 -- lib/ai/embeddingprocessor.go | 331 ---- lib/ai/embeddingprocessor_test.go | 313 ---- lib/ai/mock_embedder.go | 54 - lib/ai/model/agent.go | 436 ----- lib/ai/model/output/error.go | 60 - lib/ai/model/output/messages.go | 79 - lib/ai/model/output/parsejson.go | 48 - lib/ai/model/output/streaming.go | 83 - lib/ai/model/prompt.go | 131 -- lib/ai/model/tools/accessrequest.go | 307 ---- lib/ai/model/tools/auditquery.go | 177 -- lib/ai/model/tools/commandexec.go | 82 - lib/ai/model/tools/embedding.go | 147 -- lib/ai/model/tools/embedding_test.go | 140 -- lib/ai/model/tools/generationtool.go | 84 - lib/ai/model/tools/tool.go | 77 - lib/ai/simpleretriever.go | 146 -- lib/ai/simpleretriever_test.go | 103 -- lib/ai/testutils/http.go | 131 -- lib/ai/tokens/contants.go | 31 - lib/ai/tokens/tokencount.go | 194 -- lib/ai/tokens/tokencount_test.go | 96 - lib/assist/assist.go | 623 ------- lib/assist/assist_test.go | 209 --- lib/assist/constants.go | 37 - lib/assist/messages.go | 36 - lib/auth/assist/assistv1/service.go | 533 ------ lib/auth/assist/assistv1/test/service_test.go | 525 ------ lib/auth/auth.go | 54 - lib/auth/authclient/clt.go | 5 - lib/auth/grpcserver.go | 16 - lib/auth/helpers.go | 28 +- lib/auth/init.go | 14 - .../userpreferencesv1/service_test.go | 8 - lib/config/configuration.go | 51 +- lib/config/configuration_test.go | 59 - lib/config/fileconf.go | 34 +- lib/modules/modules.go | 4 - lib/service/proxy_settings.go | 7 - lib/service/proxy_settings_test.go | 111 -- lib/service/service.go | 62 - lib/service/servicecfg/auth.go | 4 - lib/service/servicecfg/config.go | 10 - lib/service/servicecfg/proxy.go | 4 - lib/services/assist.go | 48 - lib/services/embeddings.go | 40 - lib/services/local/assistant.go | 285 --- lib/services/local/assistant_test.go | 176 -- lib/services/local/embeddings.go | 117 -- lib/services/local/embeddings_test.go | 237 --- lib/services/local/userpreferences.go | 4 - lib/services/local/userpreferences_test.go | 93 - lib/services/local/users.go | 44 +- lib/services/presets.go | 2 - lib/services/useracl.go | 8 - .../userpreferences/userpreferences_test.go | 2 - lib/web/apiserver.go | 56 +- lib/web/apiserver_test.go | 14 - lib/web/assistant.go | 371 ---- lib/web/assistant_test.go | 203 --- lib/web/userpreferences.go | 17 - tool/tctl/common/resource_command_test.go | 1 - 86 files changed, 143 insertions(+), 11082 deletions(-) delete mode 100644 api/gen/proto/go/assist/v1/assist.pb.go delete mode 100644 api/gen/proto/go/assist/v1/assist_grpc.pb.go delete mode 100644 api/proto/teleport/assist/v1/assist.proto delete mode 100644 lib/ai/chat.go delete mode 100644 lib/ai/chat_test.go delete mode 100644 lib/ai/client.go delete mode 100644 lib/ai/client_test.go delete mode 100644 lib/ai/embedding/embedding.go delete mode 100644 lib/ai/embedding/serialization.go delete mode 100644 lib/ai/embeddingprocessor.go delete mode 100644 lib/ai/embeddingprocessor_test.go delete mode 100644 lib/ai/mock_embedder.go delete mode 100644 lib/ai/model/agent.go delete mode 100644 lib/ai/model/output/error.go delete mode 100644 lib/ai/model/output/messages.go delete mode 100644 lib/ai/model/output/parsejson.go delete mode 100644 lib/ai/model/output/streaming.go delete mode 100644 lib/ai/model/prompt.go delete mode 100644 lib/ai/model/tools/accessrequest.go delete mode 100644 lib/ai/model/tools/auditquery.go delete mode 100644 lib/ai/model/tools/commandexec.go delete mode 100644 lib/ai/model/tools/embedding.go delete mode 100644 lib/ai/model/tools/embedding_test.go delete mode 100644 lib/ai/model/tools/generationtool.go delete mode 100644 lib/ai/model/tools/tool.go delete mode 100644 lib/ai/simpleretriever.go delete mode 100644 lib/ai/simpleretriever_test.go delete mode 100644 lib/ai/testutils/http.go delete mode 100644 lib/ai/tokens/contants.go delete mode 100644 lib/ai/tokens/tokencount.go delete mode 100644 lib/ai/tokens/tokencount_test.go delete mode 100644 lib/assist/assist.go delete mode 100644 lib/assist/assist_test.go delete mode 100644 lib/assist/constants.go delete mode 100644 lib/assist/messages.go delete mode 100644 lib/auth/assist/assistv1/service.go delete mode 100644 lib/auth/assist/assistv1/test/service_test.go delete mode 100644 lib/service/proxy_settings_test.go delete mode 100644 lib/services/assist.go delete mode 100644 lib/services/embeddings.go delete mode 100644 lib/services/local/assistant.go delete mode 100644 lib/services/local/assistant_test.go delete mode 100644 lib/services/local/embeddings.go delete mode 100644 lib/services/local/embeddings_test.go delete mode 100644 lib/web/assistant.go delete mode 100644 lib/web/assistant_test.go diff --git a/.github/ISSUE_TEMPLATE/testplan.md b/.github/ISSUE_TEMPLATE/testplan.md index 1e3f33c23a2..1738b5560c5 100644 --- a/.github/ISSUE_TEMPLATE/testplan.md +++ b/.github/ISSUE_TEMPLATE/testplan.md @@ -1524,23 +1524,6 @@ Docs: [IP Pinning](https://goteleport.com/docs/access-controls/guides/ip-pinning - [ ] You can access Desktop service on leaf cluster - [ ] If you change your IP you no longer can access Desktop services. -## Assist - -Assist is not supported by `tsh` and WebUI is the only way to use it. -Assist test plan is in the core section instead of WebUI as most functionality is implemented in the core. - -- Configuration - - [ ] Assist is disabled by default (OSS, Enterprise) - - [ ] Assist can be enabled in the configuration file. - - [ ] Assist is disabled in the Cloud. - - [ ] Assist is enabled by default in the Cloud Team plan. - - [ ] Assist is always disabled when etcd is used as a backend. - -- SSH integration - - [ ] Assist icon is visible in WebUI's Terminal - - [ ] A Bash command can be generated in the above window. - - [ ] When an output is selected in the Terminal "Explain" option is available, and it generates the summary. - ## IGS: - [ ] Access Monitoring - [ ] Verify that users can run custom audit queries. diff --git a/CHANGELOG.md b/CHANGELOG.md index 51c053e58ff..0e7426513be 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,10 +12,9 @@ Opsgenie plugin users, role annotations must now contain See [the Opsgenie plugin documentation](docs/pages/access-controls/access-request-plugins/opsgenie.mdx) for setup instructions. -#### Teleport Assist chat has been removed +#### Teleport Assist has been removed -Teleport Assist chat has been removed from Teleport 16. Assist is still available -in the SSH Web Terminal and Audit Monitoring. +Teleport Assist chat has been removed from Teleport 16. #### DynamoDB permission requirements have changed diff --git a/api/client/client.go b/api/client/client.go index 6827923dfcb..74022058526 100644 --- a/api/client/client.go +++ b/api/client/client.go @@ -60,7 +60,6 @@ import ( "github.com/gravitational/teleport/api/client/userloginstate" "github.com/gravitational/teleport/api/constants" "github.com/gravitational/teleport/api/defaults" - "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" accesslistv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/accesslist/v1" accessmonitoringrulev1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/accessmonitoringrules/v1" auditlogpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/auditlog/v1" @@ -110,7 +109,6 @@ func init() { // AuthServiceClient keeps the interfaces implemented by the auth service. type AuthServiceClient struct { proto.AuthServiceClient - assist.AssistServiceClient auditlogpb.AuditLogServiceClient userpreferencespb.UserPreferencesServiceClient notificationsv1pb.NotificationServiceClient @@ -522,7 +520,6 @@ func (c *Client) dialGRPC(ctx context.Context, addr string) error { c.conn = conn c.grpc = AuthServiceClient{ AuthServiceClient: proto.NewAuthServiceClient(c.conn), - AssistServiceClient: assist.NewAssistServiceClient(c.conn), AuditLogServiceClient: auditlogpb.NewAuditLogServiceClient(c.conn), UserPreferencesServiceClient: userpreferencespb.NewUserPreferencesServiceClient(c.conn), NotificationServiceClient: notificationsv1pb.NewNotificationServiceClient(c.conn), @@ -847,12 +844,6 @@ func (c *Client) TrustClient() trustpb.TrustServiceClient { return trustpb.NewTrustServiceClient(c.conn) } -// EmbeddingClient returns an unadorned Embedding client, using the underlying -// Auth gRPC connection. -func (c *Client) EmbeddingClient() assist.AssistEmbeddingServiceClient { - return assist.NewAssistEmbeddingServiceClient(c.conn) -} - // BotServiceClient returns an unadorned client for the bot service. func (c *Client) BotServiceClient() machineidv1pb.BotServiceClient { return machineidv1pb.NewBotServiceClient(c.conn) @@ -4750,78 +4741,6 @@ func (c *Client) WatchPendingHeadlessAuthentications(ctx context.Context) (types return w, nil } -// CreateAssistantConversation creates a new conversation entry in the backend. -func (c *Client) CreateAssistantConversation(ctx context.Context, req *assist.CreateAssistantConversationRequest) (*assist.CreateAssistantConversationResponse, error) { - resp, err := c.grpc.CreateAssistantConversation(ctx, req) - if err != nil { - return nil, trace.Wrap(err) - } - - return resp, nil -} - -// GetAssistantMessages retrieves assistant messages with given conversation ID. -func (c *Client) GetAssistantMessages(ctx context.Context, req *assist.GetAssistantMessagesRequest) (*assist.GetAssistantMessagesResponse, error) { - messages, err := c.grpc.GetAssistantMessages(ctx, req) - if err != nil { - return nil, trace.Wrap(err) - } - return messages, nil -} - -// DeleteAssistantConversation deletes a conversation entry in the backend. -func (c *Client) DeleteAssistantConversation(ctx context.Context, req *assist.DeleteAssistantConversationRequest) error { - _, err := c.grpc.DeleteAssistantConversation(ctx, req) - if err != nil { - return trace.Wrap(err) - } - return nil -} - -// IsAssistEnabled returns true if the assist is enabled or not on the auth level. -func (c *Client) IsAssistEnabled(ctx context.Context) (*assist.IsAssistEnabledResponse, error) { - resp, err := c.grpc.IsAssistEnabled(ctx, &assist.IsAssistEnabledRequest{}) - if err != nil { - return nil, trace.Wrap(err) - } - return resp, nil -} - -// GetAssistantConversations returns all conversations started by a user. -func (c *Client) GetAssistantConversations(ctx context.Context, request *assist.GetAssistantConversationsRequest) (*assist.GetAssistantConversationsResponse, error) { - messages, err := c.grpc.GetAssistantConversations(ctx, request) - if err != nil { - return nil, trace.Wrap(err) - } - return messages, nil -} - -// CreateAssistantMessage saves a new conversation message. -func (c *Client) CreateAssistantMessage(ctx context.Context, in *assist.CreateAssistantMessageRequest) error { - _, err := c.grpc.CreateAssistantMessage(ctx, in) - if err != nil { - return trace.Wrap(err) - } - return nil -} - -// UpdateAssistantConversationInfo updates conversation info. -func (c *Client) UpdateAssistantConversationInfo(ctx context.Context, in *assist.UpdateAssistantConversationInfoRequest) error { - _, err := c.grpc.UpdateAssistantConversationInfo(ctx, in) - if err != nil { - return trace.Wrap(err) - } - return nil -} - -func (c *Client) GetAssistantEmbeddings(ctx context.Context, in *assist.GetAssistantEmbeddingsRequest) (*assist.GetAssistantEmbeddingsResponse, error) { - result, err := c.EmbeddingClient().GetAssistantEmbeddings(ctx, in) - if err != nil { - return nil, trace.Wrap(err) - } - return result, nil -} - // GetUserPreferences returns the user preferences for a given user. func (c *Client) GetUserPreferences(ctx context.Context, in *userpreferencespb.GetUserPreferencesRequest) (*userpreferencespb.GetUserPreferencesResponse, error) { resp, err := c.grpc.GetUserPreferences(ctx, in) diff --git a/api/client/webclient/webclient.go b/api/client/webclient/webclient.go index cdcdd271758..d2e5b4af765 100644 --- a/api/client/webclient/webclient.go +++ b/api/client/webclient/webclient.go @@ -326,8 +326,6 @@ type ProxySettings struct { // TLSRoutingEnabled indicates that proxy supports ALPN SNI server where // all proxy services are exposed on a single TLS listener (Proxy Web Listener). TLSRoutingEnabled bool `json:"tls_routing_enabled"` - // AssistEnabled is true when Teleport Assist is enabled. - AssistEnabled bool `json:"assist_enabled"` } // KubeProxySettings is kubernetes proxy settings diff --git a/api/client/webclient/webconfig.go b/api/client/webclient/webconfig.go index ce28f24865e..d711fd9ed5f 100644 --- a/api/client/webclient/webconfig.go +++ b/api/client/webclient/webconfig.go @@ -69,8 +69,6 @@ type WebConfig struct { // Eg, v13.4.3 // Only present when AutomaticUpgrades are enabled. AutomaticUpgradesTargetVersion string `json:"automaticUpgradesTargetVersion,omitempty"` - // AssistEnabled is true when Teleport Assist is enabled. - AssistEnabled bool `json:"assistEnabled"` // HideInaccessibleFeatures is true when features should be undiscoverable to users without the necessary permissions. // Usually, in order to encourage discoverability of features, we show UI elements even if the user doesn't have permission to access them, // this flag disables that behavior. diff --git a/api/defaults/defaults.go b/api/defaults/defaults.go index a4a34da93dc..88624cd12ce 100644 --- a/api/defaults/defaults.go +++ b/api/defaults/defaults.go @@ -81,10 +81,6 @@ const ( // BreakerRatioMinExecutions is the minimum number of requests before the ratio tripper // will consider examining the request pass rate BreakerRatioMinExecutions = 10 - - // AssistCommandExecutionWorkers is the number of workers that will - // execute arbitrary remote commands on servers in parallel - AssistCommandExecutionWorkers = 30 ) var ( diff --git a/api/gen/proto/go/assist/v1/assist.pb.go b/api/gen/proto/go/assist/v1/assist.pb.go deleted file mode 100644 index e702e5c557c..00000000000 --- a/api/gen/proto/go/assist/v1/assist.pb.go +++ /dev/null @@ -1,1595 +0,0 @@ -// Copyright 2023 Gravitational, Inc -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -// Code generated by protoc-gen-go. DO NOT EDIT. -// versions: -// protoc-gen-go v1.34.1 -// protoc (unknown) -// source: teleport/assist/v1/assist.proto - -package assist - -import ( - proto "github.com/gravitational/teleport/api/client/proto" - protoreflect "google.golang.org/protobuf/reflect/protoreflect" - protoimpl "google.golang.org/protobuf/runtime/protoimpl" - emptypb "google.golang.org/protobuf/types/known/emptypb" - timestamppb "google.golang.org/protobuf/types/known/timestamppb" - reflect "reflect" - sync "sync" -) - -const ( - // Verify that this generated code is sufficiently up-to-date. - _ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion) - // Verify that runtime/protoimpl is sufficiently up-to-date. - _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) -) - -// GetAssistantMessagesRequest is a request to the assistant service. -type GetAssistantMessagesRequest struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // ConversationId identifies a conversation. - // It's used to tie all messages in a one conversation. - ConversationId string `protobuf:"bytes,1,opt,name=conversation_id,json=conversationId,proto3" json:"conversation_id,omitempty"` - // username is a username of the user who sent the message. - Username string `protobuf:"bytes,2,opt,name=username,proto3" json:"username,omitempty"` -} - -func (x *GetAssistantMessagesRequest) Reset() { - *x = GetAssistantMessagesRequest{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[0] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *GetAssistantMessagesRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*GetAssistantMessagesRequest) ProtoMessage() {} - -func (x *GetAssistantMessagesRequest) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[0] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use GetAssistantMessagesRequest.ProtoReflect.Descriptor instead. -func (*GetAssistantMessagesRequest) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{0} -} - -func (x *GetAssistantMessagesRequest) GetConversationId() string { - if x != nil { - return x.ConversationId - } - return "" -} - -func (x *GetAssistantMessagesRequest) GetUsername() string { - if x != nil { - return x.Username - } - return "" -} - -// AssistantMessage is a message sent to the assistant service. The conversation -// must be created first. -type AssistantMessage struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // type is a type of message. It can be Chat response/query or a command to run. - Type string `protobuf:"bytes,1,opt,name=type,proto3" json:"type,omitempty"` - // CreatedTime is the time when the event occurred. - CreatedTime *timestamppb.Timestamp `protobuf:"bytes,2,opt,name=created_time,json=createdTime,proto3" json:"created_time,omitempty"` - // payload is a JSON message - Payload string `protobuf:"bytes,3,opt,name=payload,proto3" json:"payload,omitempty"` -} - -func (x *AssistantMessage) Reset() { - *x = AssistantMessage{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[1] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *AssistantMessage) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*AssistantMessage) ProtoMessage() {} - -func (x *AssistantMessage) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[1] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use AssistantMessage.ProtoReflect.Descriptor instead. -func (*AssistantMessage) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{1} -} - -func (x *AssistantMessage) GetType() string { - if x != nil { - return x.Type - } - return "" -} - -func (x *AssistantMessage) GetCreatedTime() *timestamppb.Timestamp { - if x != nil { - return x.CreatedTime - } - return nil -} - -func (x *AssistantMessage) GetPayload() string { - if x != nil { - return x.Payload - } - return "" -} - -// CreateAssistantMessageRequest is a request to the assistant service. -type CreateAssistantMessageRequest struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // message is a message sent to the assistant service. - Message *AssistantMessage `protobuf:"bytes,1,opt,name=message,proto3" json:"message,omitempty"` - // ConversationId is used to tie all messages into a conversation. - ConversationId string `protobuf:"bytes,2,opt,name=conversation_id,json=conversationId,proto3" json:"conversation_id,omitempty"` - // username is a username of the user who sent the message. - Username string `protobuf:"bytes,3,opt,name=username,proto3" json:"username,omitempty"` -} - -func (x *CreateAssistantMessageRequest) Reset() { - *x = CreateAssistantMessageRequest{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[2] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *CreateAssistantMessageRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*CreateAssistantMessageRequest) ProtoMessage() {} - -func (x *CreateAssistantMessageRequest) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[2] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use CreateAssistantMessageRequest.ProtoReflect.Descriptor instead. -func (*CreateAssistantMessageRequest) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{2} -} - -func (x *CreateAssistantMessageRequest) GetMessage() *AssistantMessage { - if x != nil { - return x.Message - } - return nil -} - -func (x *CreateAssistantMessageRequest) GetConversationId() string { - if x != nil { - return x.ConversationId - } - return "" -} - -func (x *CreateAssistantMessageRequest) GetUsername() string { - if x != nil { - return x.Username - } - return "" -} - -// GetAssistantMessagesResponse is a response from the assistant service. -type GetAssistantMessagesResponse struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // messages is a list of messages. - Messages []*AssistantMessage `protobuf:"bytes,1,rep,name=messages,proto3" json:"messages,omitempty"` -} - -func (x *GetAssistantMessagesResponse) Reset() { - *x = GetAssistantMessagesResponse{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[3] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *GetAssistantMessagesResponse) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*GetAssistantMessagesResponse) ProtoMessage() {} - -func (x *GetAssistantMessagesResponse) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[3] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use GetAssistantMessagesResponse.ProtoReflect.Descriptor instead. -func (*GetAssistantMessagesResponse) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{3} -} - -func (x *GetAssistantMessagesResponse) GetMessages() []*AssistantMessage { - if x != nil { - return x.Messages - } - return nil -} - -// GetAssistantConversationsRequest is a request to get a list of conversations. -type GetAssistantConversationsRequest struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // username is a username of the user who created the conversation. - Username string `protobuf:"bytes,1,opt,name=username,proto3" json:"username,omitempty"` -} - -func (x *GetAssistantConversationsRequest) Reset() { - *x = GetAssistantConversationsRequest{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[4] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *GetAssistantConversationsRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*GetAssistantConversationsRequest) ProtoMessage() {} - -func (x *GetAssistantConversationsRequest) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[4] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use GetAssistantConversationsRequest.ProtoReflect.Descriptor instead. -func (*GetAssistantConversationsRequest) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{4} -} - -func (x *GetAssistantConversationsRequest) GetUsername() string { - if x != nil { - return x.Username - } - return "" -} - -// ConversationInfo is a conversation info. It contains a conversation -// information like ID, title, created time. -type ConversationInfo struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // id is a unique conversation ID. - Id string `protobuf:"bytes,1,opt,name=id,proto3" json:"id,omitempty"` - // title is a title of the conversation. - Title string `protobuf:"bytes,2,opt,name=title,proto3" json:"title,omitempty"` - // createdTime is the time when the conversation was created. - CreatedTime *timestamppb.Timestamp `protobuf:"bytes,3,opt,name=created_time,json=createdTime,proto3" json:"created_time,omitempty"` -} - -func (x *ConversationInfo) Reset() { - *x = ConversationInfo{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[5] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *ConversationInfo) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*ConversationInfo) ProtoMessage() {} - -func (x *ConversationInfo) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[5] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use ConversationInfo.ProtoReflect.Descriptor instead. -func (*ConversationInfo) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{5} -} - -func (x *ConversationInfo) GetId() string { - if x != nil { - return x.Id - } - return "" -} - -func (x *ConversationInfo) GetTitle() string { - if x != nil { - return x.Title - } - return "" -} - -func (x *ConversationInfo) GetCreatedTime() *timestamppb.Timestamp { - if x != nil { - return x.CreatedTime - } - return nil -} - -// GetAssistantConversationsResponse is a response from the assistant service. -type GetAssistantConversationsResponse struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // conversations is a list of conversations. - Conversations []*ConversationInfo `protobuf:"bytes,1,rep,name=conversations,proto3" json:"conversations,omitempty"` -} - -func (x *GetAssistantConversationsResponse) Reset() { - *x = GetAssistantConversationsResponse{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[6] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *GetAssistantConversationsResponse) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*GetAssistantConversationsResponse) ProtoMessage() {} - -func (x *GetAssistantConversationsResponse) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[6] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use GetAssistantConversationsResponse.ProtoReflect.Descriptor instead. -func (*GetAssistantConversationsResponse) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{6} -} - -func (x *GetAssistantConversationsResponse) GetConversations() []*ConversationInfo { - if x != nil { - return x.Conversations - } - return nil -} - -// CreateAssistantConversationRequest is a request to create a new conversation. -type CreateAssistantConversationRequest struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // username is a username of the user who created the conversation. - Username string `protobuf:"bytes,1,opt,name=username,proto3" json:"username,omitempty"` - // createdTime is the time when the conversation was created. - CreatedTime *timestamppb.Timestamp `protobuf:"bytes,2,opt,name=created_time,json=createdTime,proto3" json:"created_time,omitempty"` -} - -func (x *CreateAssistantConversationRequest) Reset() { - *x = CreateAssistantConversationRequest{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[7] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *CreateAssistantConversationRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*CreateAssistantConversationRequest) ProtoMessage() {} - -func (x *CreateAssistantConversationRequest) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[7] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use CreateAssistantConversationRequest.ProtoReflect.Descriptor instead. -func (*CreateAssistantConversationRequest) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{7} -} - -func (x *CreateAssistantConversationRequest) GetUsername() string { - if x != nil { - return x.Username - } - return "" -} - -func (x *CreateAssistantConversationRequest) GetCreatedTime() *timestamppb.Timestamp { - if x != nil { - return x.CreatedTime - } - return nil -} - -// CreateAssistantConversationResponse is a response from the assistant service. -type CreateAssistantConversationResponse struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // id is a unique conversation ID. - Id string `protobuf:"bytes,1,opt,name=id,proto3" json:"id,omitempty"` -} - -func (x *CreateAssistantConversationResponse) Reset() { - *x = CreateAssistantConversationResponse{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[8] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *CreateAssistantConversationResponse) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*CreateAssistantConversationResponse) ProtoMessage() {} - -func (x *CreateAssistantConversationResponse) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[8] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use CreateAssistantConversationResponse.ProtoReflect.Descriptor instead. -func (*CreateAssistantConversationResponse) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{8} -} - -func (x *CreateAssistantConversationResponse) GetId() string { - if x != nil { - return x.Id - } - return "" -} - -// UpdateAssistantConversationInfoRequest is a request to update the conversation info. -type UpdateAssistantConversationInfoRequest struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // conversationId is a unique conversation ID. - ConversationId string `protobuf:"bytes,1,opt,name=conversation_id,json=conversationId,proto3" json:"conversation_id,omitempty"` - // username is a username of the user who created the conversation. - Username string `protobuf:"bytes,2,opt,name=username,proto3" json:"username,omitempty"` - // title is a title of the conversation. - Title string `protobuf:"bytes,3,opt,name=title,proto3" json:"title,omitempty"` -} - -func (x *UpdateAssistantConversationInfoRequest) Reset() { - *x = UpdateAssistantConversationInfoRequest{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[9] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *UpdateAssistantConversationInfoRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*UpdateAssistantConversationInfoRequest) ProtoMessage() {} - -func (x *UpdateAssistantConversationInfoRequest) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[9] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use UpdateAssistantConversationInfoRequest.ProtoReflect.Descriptor instead. -func (*UpdateAssistantConversationInfoRequest) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{9} -} - -func (x *UpdateAssistantConversationInfoRequest) GetConversationId() string { - if x != nil { - return x.ConversationId - } - return "" -} - -func (x *UpdateAssistantConversationInfoRequest) GetUsername() string { - if x != nil { - return x.Username - } - return "" -} - -func (x *UpdateAssistantConversationInfoRequest) GetTitle() string { - if x != nil { - return x.Title - } - return "" -} - -// IsAssistEnabledRequest is a request to the assistant service on if assist is enabled or not. -type IsAssistEnabledRequest struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields -} - -func (x *IsAssistEnabledRequest) Reset() { - *x = IsAssistEnabledRequest{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[10] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *IsAssistEnabledRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*IsAssistEnabledRequest) ProtoMessage() {} - -func (x *IsAssistEnabledRequest) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[10] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use IsAssistEnabledRequest.ProtoReflect.Descriptor instead. -func (*IsAssistEnabledRequest) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{10} -} - -// IsAssistEnabledResponse is a response from the assistant service on if assist is enabled or not. -type IsAssistEnabledResponse struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // enabled is true if the assist is enabled or not on the auth level. - Enabled bool `protobuf:"varint,1,opt,name=enabled,proto3" json:"enabled,omitempty"` -} - -func (x *IsAssistEnabledResponse) Reset() { - *x = IsAssistEnabledResponse{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[11] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *IsAssistEnabledResponse) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*IsAssistEnabledResponse) ProtoMessage() {} - -func (x *IsAssistEnabledResponse) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[11] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use IsAssistEnabledResponse.ProtoReflect.Descriptor instead. -func (*IsAssistEnabledResponse) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{11} -} - -func (x *IsAssistEnabledResponse) GetEnabled() bool { - if x != nil { - return x.Enabled - } - return false -} - -// DeleteAssistantConversationRequest is a request to delete the conversation. -type DeleteAssistantConversationRequest struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // conversationId is a unique conversation ID. - ConversationId string `protobuf:"bytes,1,opt,name=conversation_id,json=conversationId,proto3" json:"conversation_id,omitempty"` - // username is a username of the user who created the conversation. - Username string `protobuf:"bytes,2,opt,name=username,proto3" json:"username,omitempty"` -} - -func (x *DeleteAssistantConversationRequest) Reset() { - *x = DeleteAssistantConversationRequest{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[12] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *DeleteAssistantConversationRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*DeleteAssistantConversationRequest) ProtoMessage() {} - -func (x *DeleteAssistantConversationRequest) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[12] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use DeleteAssistantConversationRequest.ProtoReflect.Descriptor instead. -func (*DeleteAssistantConversationRequest) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{12} -} - -func (x *DeleteAssistantConversationRequest) GetConversationId() string { - if x != nil { - return x.ConversationId - } - return "" -} - -func (x *DeleteAssistantConversationRequest) GetUsername() string { - if x != nil { - return x.Username - } - return "" -} - -// GetAssistantEmbeddingsRequest is a request to get embeddings. -type GetAssistantEmbeddingsRequest struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // username is a username of the user who requested the embeddings. - Username string `protobuf:"bytes,1,opt,name=username,proto3" json:"username,omitempty"` - // query is the query used for similarity search. - Query string `protobuf:"bytes,2,opt,name=query,proto3" json:"query,omitempty"` - // limit is the number of embeddings to return (also known as k). - Limit uint32 `protobuf:"varint,3,opt,name=limit,proto3" json:"limit,omitempty"` - // kind is the kind of embeddings to return (ex, node). - Kind string `protobuf:"bytes,4,opt,name=kind,proto3" json:"kind,omitempty"` -} - -func (x *GetAssistantEmbeddingsRequest) Reset() { - *x = GetAssistantEmbeddingsRequest{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[13] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *GetAssistantEmbeddingsRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*GetAssistantEmbeddingsRequest) ProtoMessage() {} - -func (x *GetAssistantEmbeddingsRequest) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[13] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use GetAssistantEmbeddingsRequest.ProtoReflect.Descriptor instead. -func (*GetAssistantEmbeddingsRequest) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{13} -} - -func (x *GetAssistantEmbeddingsRequest) GetUsername() string { - if x != nil { - return x.Username - } - return "" -} - -func (x *GetAssistantEmbeddingsRequest) GetQuery() string { - if x != nil { - return x.Query - } - return "" -} - -func (x *GetAssistantEmbeddingsRequest) GetLimit() uint32 { - if x != nil { - return x.Limit - } - return 0 -} - -func (x *GetAssistantEmbeddingsRequest) GetKind() string { - if x != nil { - return x.Kind - } - return "" -} - -// EmbeddingDocument is a document with an embedding. -type EmbeddedDocument struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // id is the id of the document. - Id string `protobuf:"bytes,1,opt,name=id,proto3" json:"id,omitempty"` - // content is the content of the document. - Content string `protobuf:"bytes,2,opt,name=content,proto3" json:"content,omitempty"` - // similarityScore is the similarity score of the document. - SimilarityScore float32 `protobuf:"fixed32,3,opt,name=similarity_score,json=similarityScore,proto3" json:"similarity_score,omitempty"` -} - -func (x *EmbeddedDocument) Reset() { - *x = EmbeddedDocument{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[14] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *EmbeddedDocument) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*EmbeddedDocument) ProtoMessage() {} - -func (x *EmbeddedDocument) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[14] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use EmbeddedDocument.ProtoReflect.Descriptor instead. -func (*EmbeddedDocument) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{14} -} - -func (x *EmbeddedDocument) GetId() string { - if x != nil { - return x.Id - } - return "" -} - -func (x *EmbeddedDocument) GetContent() string { - if x != nil { - return x.Content - } - return "" -} - -func (x *EmbeddedDocument) GetSimilarityScore() float32 { - if x != nil { - return x.SimilarityScore - } - return 0 -} - -// GetAssistantEmbeddingsResponse is a response from the assistant service. -type GetAssistantEmbeddingsResponse struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // embeddings is the list of embeddings. - // The list is sorted by similarity score in descending order. - Embeddings []*EmbeddedDocument `protobuf:"bytes,1,rep,name=embeddings,proto3" json:"embeddings,omitempty"` -} - -func (x *GetAssistantEmbeddingsResponse) Reset() { - *x = GetAssistantEmbeddingsResponse{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[15] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *GetAssistantEmbeddingsResponse) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*GetAssistantEmbeddingsResponse) ProtoMessage() {} - -func (x *GetAssistantEmbeddingsResponse) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[15] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use GetAssistantEmbeddingsResponse.ProtoReflect.Descriptor instead. -func (*GetAssistantEmbeddingsResponse) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{15} -} - -func (x *GetAssistantEmbeddingsResponse) GetEmbeddings() []*EmbeddedDocument { - if x != nil { - return x.Embeddings - } - return nil -} - -// SearchUnifiedResourcesRequest is a request to search for one or more resource kinds using similiarity search. -type SearchUnifiedResourcesRequest struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // query is the query used for similarity search. - Query string `protobuf:"bytes,1,opt,name=query,proto3" json:"query,omitempty"` - // limit is the number of embeddings to return (also known as k). - Limit int32 `protobuf:"varint,2,opt,name=limit,proto3" json:"limit,omitempty"` - // kinds is the kind of embeddings to return (ex, node). Returns all supported kinds if empty. - Kinds []string `protobuf:"bytes,3,rep,name=kinds,proto3" json:"kinds,omitempty"` -} - -func (x *SearchUnifiedResourcesRequest) Reset() { - *x = SearchUnifiedResourcesRequest{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[16] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *SearchUnifiedResourcesRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*SearchUnifiedResourcesRequest) ProtoMessage() {} - -func (x *SearchUnifiedResourcesRequest) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[16] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use SearchUnifiedResourcesRequest.ProtoReflect.Descriptor instead. -func (*SearchUnifiedResourcesRequest) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{16} -} - -func (x *SearchUnifiedResourcesRequest) GetQuery() string { - if x != nil { - return x.Query - } - return "" -} - -func (x *SearchUnifiedResourcesRequest) GetLimit() int32 { - if x != nil { - return x.Limit - } - return 0 -} - -func (x *SearchUnifiedResourcesRequest) GetKinds() []string { - if x != nil { - return x.Kinds - } - return nil -} - -// SearchUnifiedResourcesResponse is a response from the assistant service with a similarity-ordered list of resources. -type SearchUnifiedResourcesResponse struct { - state protoimpl.MessageState - sizeCache protoimpl.SizeCache - unknownFields protoimpl.UnknownFields - - // resources is the list of resources. - Resources []*proto.PaginatedResource `protobuf:"bytes,1,rep,name=resources,proto3" json:"resources,omitempty"` -} - -func (x *SearchUnifiedResourcesResponse) Reset() { - *x = SearchUnifiedResourcesResponse{} - if protoimpl.UnsafeEnabled { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[17] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) - } -} - -func (x *SearchUnifiedResourcesResponse) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*SearchUnifiedResourcesResponse) ProtoMessage() {} - -func (x *SearchUnifiedResourcesResponse) ProtoReflect() protoreflect.Message { - mi := &file_teleport_assist_v1_assist_proto_msgTypes[17] - if protoimpl.UnsafeEnabled && x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use SearchUnifiedResourcesResponse.ProtoReflect.Descriptor instead. -func (*SearchUnifiedResourcesResponse) Descriptor() ([]byte, []int) { - return file_teleport_assist_v1_assist_proto_rawDescGZIP(), []int{17} -} - -func (x *SearchUnifiedResourcesResponse) GetResources() []*proto.PaginatedResource { - if x != nil { - return x.Resources - } - return nil -} - -var File_teleport_assist_v1_assist_proto protoreflect.FileDescriptor - -var file_teleport_assist_v1_assist_proto_rawDesc = []byte{ - 0x0a, 0x1f, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2f, 0x61, 0x73, 0x73, 0x69, 0x73, - 0x74, 0x2f, 0x76, 0x31, 0x2f, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x70, 0x72, 0x6f, 0x74, - 0x6f, 0x12, 0x12, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, - 0x73, 0x74, 0x2e, 0x76, 0x31, 0x1a, 0x1b, 0x67, 0x6f, 0x6f, 0x67, 0x6c, 0x65, 0x2f, 0x70, 0x72, - 0x6f, 0x74, 0x6f, 0x62, 0x75, 0x66, 0x2f, 0x65, 0x6d, 0x70, 0x74, 0x79, 0x2e, 0x70, 0x72, 0x6f, - 0x74, 0x6f, 0x1a, 0x1f, 0x67, 0x6f, 0x6f, 0x67, 0x6c, 0x65, 0x2f, 0x70, 0x72, 0x6f, 0x74, 0x6f, - 0x62, 0x75, 0x66, 0x2f, 0x74, 0x69, 0x6d, 0x65, 0x73, 0x74, 0x61, 0x6d, 0x70, 0x2e, 0x70, 0x72, - 0x6f, 0x74, 0x6f, 0x1a, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2f, 0x6c, 0x65, - 0x67, 0x61, 0x63, 0x79, 0x2f, 0x63, 0x6c, 0x69, 0x65, 0x6e, 0x74, 0x2f, 0x70, 0x72, 0x6f, 0x74, - 0x6f, 0x2f, 0x61, 0x75, 0x74, 0x68, 0x73, 0x65, 0x72, 0x76, 0x69, 0x63, 0x65, 0x2e, 0x70, 0x72, - 0x6f, 0x74, 0x6f, 0x22, 0x62, 0x0a, 0x1b, 0x47, 0x65, 0x74, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, - 0x61, 0x6e, 0x74, 0x4d, 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, - 0x73, 0x74, 0x12, 0x27, 0x0a, 0x0f, 0x63, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, - 0x6f, 0x6e, 0x5f, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x0e, 0x63, 0x6f, 0x6e, - 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x49, 0x64, 0x12, 0x1a, 0x0a, 0x08, 0x75, - 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, 0x65, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x08, 0x75, - 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, 0x65, 0x22, 0x7f, 0x0a, 0x10, 0x41, 0x73, 0x73, 0x69, 0x73, - 0x74, 0x61, 0x6e, 0x74, 0x4d, 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, 0x12, 0x12, 0x0a, 0x04, 0x74, - 0x79, 0x70, 0x65, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x04, 0x74, 0x79, 0x70, 0x65, 0x12, - 0x3d, 0x0a, 0x0c, 0x63, 0x72, 0x65, 0x61, 0x74, 0x65, 0x64, 0x5f, 0x74, 0x69, 0x6d, 0x65, 0x18, - 0x02, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x1a, 0x2e, 0x67, 0x6f, 0x6f, 0x67, 0x6c, 0x65, 0x2e, 0x70, - 0x72, 0x6f, 0x74, 0x6f, 0x62, 0x75, 0x66, 0x2e, 0x54, 0x69, 0x6d, 0x65, 0x73, 0x74, 0x61, 0x6d, - 0x70, 0x52, 0x0b, 0x63, 0x72, 0x65, 0x61, 0x74, 0x65, 0x64, 0x54, 0x69, 0x6d, 0x65, 0x12, 0x18, - 0x0a, 0x07, 0x70, 0x61, 0x79, 0x6c, 0x6f, 0x61, 0x64, 0x18, 0x03, 0x20, 0x01, 0x28, 0x09, 0x52, - 0x07, 0x70, 0x61, 0x79, 0x6c, 0x6f, 0x61, 0x64, 0x22, 0xa4, 0x01, 0x0a, 0x1d, 0x43, 0x72, 0x65, - 0x61, 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x4d, 0x65, 0x73, 0x73, - 0x61, 0x67, 0x65, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x3e, 0x0a, 0x07, 0x6d, 0x65, - 0x73, 0x73, 0x61, 0x67, 0x65, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x24, 0x2e, 0x74, 0x65, - 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, - 0x2e, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x4d, 0x65, 0x73, 0x73, 0x61, 0x67, - 0x65, 0x52, 0x07, 0x6d, 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, 0x12, 0x27, 0x0a, 0x0f, 0x63, 0x6f, - 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x5f, 0x69, 0x64, 0x18, 0x02, 0x20, - 0x01, 0x28, 0x09, 0x52, 0x0e, 0x63, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, - 0x6e, 0x49, 0x64, 0x12, 0x1a, 0x0a, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, 0x65, 0x18, - 0x03, 0x20, 0x01, 0x28, 0x09, 0x52, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, 0x65, 0x22, - 0x60, 0x0a, 0x1c, 0x47, 0x65, 0x74, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x4d, - 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, - 0x40, 0x0a, 0x08, 0x6d, 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, 0x73, 0x18, 0x01, 0x20, 0x03, 0x28, - 0x0b, 0x32, 0x24, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, - 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, - 0x4d, 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, 0x52, 0x08, 0x6d, 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, - 0x73, 0x22, 0x3e, 0x0a, 0x20, 0x47, 0x65, 0x74, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, - 0x74, 0x43, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x73, 0x52, 0x65, - 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x1a, 0x0a, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, - 0x65, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, - 0x65, 0x22, 0x77, 0x0a, 0x10, 0x43, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, - 0x6e, 0x49, 0x6e, 0x66, 0x6f, 0x12, 0x0e, 0x0a, 0x02, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, - 0x09, 0x52, 0x02, 0x69, 0x64, 0x12, 0x14, 0x0a, 0x05, 0x74, 0x69, 0x74, 0x6c, 0x65, 0x18, 0x02, - 0x20, 0x01, 0x28, 0x09, 0x52, 0x05, 0x74, 0x69, 0x74, 0x6c, 0x65, 0x12, 0x3d, 0x0a, 0x0c, 0x63, - 0x72, 0x65, 0x61, 0x74, 0x65, 0x64, 0x5f, 0x74, 0x69, 0x6d, 0x65, 0x18, 0x03, 0x20, 0x01, 0x28, - 0x0b, 0x32, 0x1a, 0x2e, 0x67, 0x6f, 0x6f, 0x67, 0x6c, 0x65, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, - 0x62, 0x75, 0x66, 0x2e, 0x54, 0x69, 0x6d, 0x65, 0x73, 0x74, 0x61, 0x6d, 0x70, 0x52, 0x0b, 0x63, - 0x72, 0x65, 0x61, 0x74, 0x65, 0x64, 0x54, 0x69, 0x6d, 0x65, 0x22, 0x6f, 0x0a, 0x21, 0x47, 0x65, - 0x74, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x43, 0x6f, 0x6e, 0x76, 0x65, 0x72, - 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, - 0x4a, 0x0a, 0x0d, 0x63, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x73, - 0x18, 0x01, 0x20, 0x03, 0x28, 0x0b, 0x32, 0x24, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, - 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x43, 0x6f, 0x6e, 0x76, - 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x49, 0x6e, 0x66, 0x6f, 0x52, 0x0d, 0x63, 0x6f, - 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x73, 0x22, 0x7f, 0x0a, 0x22, 0x43, - 0x72, 0x65, 0x61, 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x43, 0x6f, - 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, - 0x74, 0x12, 0x1a, 0x0a, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, 0x65, 0x18, 0x01, 0x20, - 0x01, 0x28, 0x09, 0x52, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, 0x65, 0x12, 0x3d, 0x0a, - 0x0c, 0x63, 0x72, 0x65, 0x61, 0x74, 0x65, 0x64, 0x5f, 0x74, 0x69, 0x6d, 0x65, 0x18, 0x02, 0x20, - 0x01, 0x28, 0x0b, 0x32, 0x1a, 0x2e, 0x67, 0x6f, 0x6f, 0x67, 0x6c, 0x65, 0x2e, 0x70, 0x72, 0x6f, - 0x74, 0x6f, 0x62, 0x75, 0x66, 0x2e, 0x54, 0x69, 0x6d, 0x65, 0x73, 0x74, 0x61, 0x6d, 0x70, 0x52, - 0x0b, 0x63, 0x72, 0x65, 0x61, 0x74, 0x65, 0x64, 0x54, 0x69, 0x6d, 0x65, 0x22, 0x35, 0x0a, 0x23, - 0x43, 0x72, 0x65, 0x61, 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x43, - 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x73, 0x70, 0x6f, - 0x6e, 0x73, 0x65, 0x12, 0x0e, 0x0a, 0x02, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, - 0x02, 0x69, 0x64, 0x22, 0x83, 0x01, 0x0a, 0x26, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x41, 0x73, - 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x43, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, - 0x69, 0x6f, 0x6e, 0x49, 0x6e, 0x66, 0x6f, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x27, - 0x0a, 0x0f, 0x63, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x5f, 0x69, - 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x0e, 0x63, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, - 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x49, 0x64, 0x12, 0x1a, 0x0a, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, - 0x61, 0x6d, 0x65, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, - 0x61, 0x6d, 0x65, 0x12, 0x14, 0x0a, 0x05, 0x74, 0x69, 0x74, 0x6c, 0x65, 0x18, 0x03, 0x20, 0x01, - 0x28, 0x09, 0x52, 0x05, 0x74, 0x69, 0x74, 0x6c, 0x65, 0x22, 0x18, 0x0a, 0x16, 0x49, 0x73, 0x41, - 0x73, 0x73, 0x69, 0x73, 0x74, 0x45, 0x6e, 0x61, 0x62, 0x6c, 0x65, 0x64, 0x52, 0x65, 0x71, 0x75, - 0x65, 0x73, 0x74, 0x22, 0x33, 0x0a, 0x17, 0x49, 0x73, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x45, - 0x6e, 0x61, 0x62, 0x6c, 0x65, 0x64, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x18, - 0x0a, 0x07, 0x65, 0x6e, 0x61, 0x62, 0x6c, 0x65, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x08, 0x52, - 0x07, 0x65, 0x6e, 0x61, 0x62, 0x6c, 0x65, 0x64, 0x22, 0x69, 0x0a, 0x22, 0x44, 0x65, 0x6c, 0x65, - 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x43, 0x6f, 0x6e, 0x76, 0x65, - 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x27, - 0x0a, 0x0f, 0x63, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x5f, 0x69, - 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x0e, 0x63, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, - 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x49, 0x64, 0x12, 0x1a, 0x0a, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, - 0x61, 0x6d, 0x65, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, - 0x61, 0x6d, 0x65, 0x22, 0x7b, 0x0a, 0x1d, 0x47, 0x65, 0x74, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, - 0x61, 0x6e, 0x74, 0x45, 0x6d, 0x62, 0x65, 0x64, 0x64, 0x69, 0x6e, 0x67, 0x73, 0x52, 0x65, 0x71, - 0x75, 0x65, 0x73, 0x74, 0x12, 0x1a, 0x0a, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, 0x65, - 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, 0x65, - 0x12, 0x14, 0x0a, 0x05, 0x71, 0x75, 0x65, 0x72, 0x79, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, - 0x05, 0x71, 0x75, 0x65, 0x72, 0x79, 0x12, 0x14, 0x0a, 0x05, 0x6c, 0x69, 0x6d, 0x69, 0x74, 0x18, - 0x03, 0x20, 0x01, 0x28, 0x0d, 0x52, 0x05, 0x6c, 0x69, 0x6d, 0x69, 0x74, 0x12, 0x12, 0x0a, 0x04, - 0x6b, 0x69, 0x6e, 0x64, 0x18, 0x04, 0x20, 0x01, 0x28, 0x09, 0x52, 0x04, 0x6b, 0x69, 0x6e, 0x64, - 0x22, 0x67, 0x0a, 0x10, 0x45, 0x6d, 0x62, 0x65, 0x64, 0x64, 0x65, 0x64, 0x44, 0x6f, 0x63, 0x75, - 0x6d, 0x65, 0x6e, 0x74, 0x12, 0x0e, 0x0a, 0x02, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, - 0x52, 0x02, 0x69, 0x64, 0x12, 0x18, 0x0a, 0x07, 0x63, 0x6f, 0x6e, 0x74, 0x65, 0x6e, 0x74, 0x18, - 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x07, 0x63, 0x6f, 0x6e, 0x74, 0x65, 0x6e, 0x74, 0x12, 0x29, - 0x0a, 0x10, 0x73, 0x69, 0x6d, 0x69, 0x6c, 0x61, 0x72, 0x69, 0x74, 0x79, 0x5f, 0x73, 0x63, 0x6f, - 0x72, 0x65, 0x18, 0x03, 0x20, 0x01, 0x28, 0x02, 0x52, 0x0f, 0x73, 0x69, 0x6d, 0x69, 0x6c, 0x61, - 0x72, 0x69, 0x74, 0x79, 0x53, 0x63, 0x6f, 0x72, 0x65, 0x22, 0x66, 0x0a, 0x1e, 0x47, 0x65, 0x74, - 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x45, 0x6d, 0x62, 0x65, 0x64, 0x64, 0x69, - 0x6e, 0x67, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x44, 0x0a, 0x0a, 0x65, - 0x6d, 0x62, 0x65, 0x64, 0x64, 0x69, 0x6e, 0x67, 0x73, 0x18, 0x01, 0x20, 0x03, 0x28, 0x0b, 0x32, - 0x24, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, - 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x45, 0x6d, 0x62, 0x65, 0x64, 0x64, 0x65, 0x64, 0x44, 0x6f, 0x63, - 0x75, 0x6d, 0x65, 0x6e, 0x74, 0x52, 0x0a, 0x65, 0x6d, 0x62, 0x65, 0x64, 0x64, 0x69, 0x6e, 0x67, - 0x73, 0x22, 0x61, 0x0a, 0x1d, 0x53, 0x65, 0x61, 0x72, 0x63, 0x68, 0x55, 0x6e, 0x69, 0x66, 0x69, - 0x65, 0x64, 0x52, 0x65, 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, - 0x73, 0x74, 0x12, 0x14, 0x0a, 0x05, 0x71, 0x75, 0x65, 0x72, 0x79, 0x18, 0x01, 0x20, 0x01, 0x28, - 0x09, 0x52, 0x05, 0x71, 0x75, 0x65, 0x72, 0x79, 0x12, 0x14, 0x0a, 0x05, 0x6c, 0x69, 0x6d, 0x69, - 0x74, 0x18, 0x02, 0x20, 0x01, 0x28, 0x05, 0x52, 0x05, 0x6c, 0x69, 0x6d, 0x69, 0x74, 0x12, 0x14, - 0x0a, 0x05, 0x6b, 0x69, 0x6e, 0x64, 0x73, 0x18, 0x03, 0x20, 0x03, 0x28, 0x09, 0x52, 0x05, 0x6b, - 0x69, 0x6e, 0x64, 0x73, 0x22, 0x58, 0x0a, 0x1e, 0x53, 0x65, 0x61, 0x72, 0x63, 0x68, 0x55, 0x6e, - 0x69, 0x66, 0x69, 0x65, 0x64, 0x52, 0x65, 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x73, 0x52, 0x65, - 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x36, 0x0a, 0x09, 0x72, 0x65, 0x73, 0x6f, 0x75, 0x72, - 0x63, 0x65, 0x73, 0x18, 0x01, 0x20, 0x03, 0x28, 0x0b, 0x32, 0x18, 0x2e, 0x70, 0x72, 0x6f, 0x74, - 0x6f, 0x2e, 0x50, 0x61, 0x67, 0x69, 0x6e, 0x61, 0x74, 0x65, 0x64, 0x52, 0x65, 0x73, 0x6f, 0x75, - 0x72, 0x63, 0x65, 0x52, 0x09, 0x72, 0x65, 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x73, 0x32, 0xde, - 0x07, 0x0a, 0x0d, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x53, 0x65, 0x72, 0x76, 0x69, 0x63, 0x65, - 0x12, 0x8e, 0x01, 0x0a, 0x1b, 0x43, 0x72, 0x65, 0x61, 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, 0x73, - 0x74, 0x61, 0x6e, 0x74, 0x43, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, - 0x12, 0x36, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, - 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x43, 0x72, 0x65, 0x61, 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, - 0x73, 0x74, 0x61, 0x6e, 0x74, 0x43, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, - 0x6e, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x37, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, - 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x43, 0x72, - 0x65, 0x61, 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x43, 0x6f, 0x6e, - 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, - 0x65, 0x12, 0x88, 0x01, 0x0a, 0x19, 0x47, 0x65, 0x74, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, - 0x6e, 0x74, 0x43, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x73, 0x12, - 0x34, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, - 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x47, 0x65, 0x74, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, - 0x74, 0x43, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x73, 0x52, 0x65, - 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x35, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, - 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x47, 0x65, 0x74, 0x41, 0x73, - 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x43, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, - 0x69, 0x6f, 0x6e, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x6d, 0x0a, 0x1b, - 0x44, 0x65, 0x6c, 0x65, 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x43, - 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x12, 0x36, 0x2e, 0x74, 0x65, - 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, - 0x2e, 0x44, 0x65, 0x6c, 0x65, 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, - 0x43, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x71, 0x75, - 0x65, 0x73, 0x74, 0x1a, 0x16, 0x2e, 0x67, 0x6f, 0x6f, 0x67, 0x6c, 0x65, 0x2e, 0x70, 0x72, 0x6f, - 0x74, 0x6f, 0x62, 0x75, 0x66, 0x2e, 0x45, 0x6d, 0x70, 0x74, 0x79, 0x12, 0x79, 0x0a, 0x14, 0x47, - 0x65, 0x74, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x4d, 0x65, 0x73, 0x73, 0x61, - 0x67, 0x65, 0x73, 0x12, 0x2f, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, - 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x47, 0x65, 0x74, 0x41, 0x73, 0x73, 0x69, - 0x73, 0x74, 0x61, 0x6e, 0x74, 0x4d, 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, 0x73, 0x52, 0x65, 0x71, - 0x75, 0x65, 0x73, 0x74, 0x1a, 0x30, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, - 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x47, 0x65, 0x74, 0x41, 0x73, 0x73, - 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x4d, 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, 0x73, 0x52, 0x65, - 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x63, 0x0a, 0x16, 0x43, 0x72, 0x65, 0x61, 0x74, 0x65, - 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x4d, 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, - 0x12, 0x31, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, - 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x43, 0x72, 0x65, 0x61, 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, - 0x73, 0x74, 0x61, 0x6e, 0x74, 0x4d, 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, 0x52, 0x65, 0x71, 0x75, - 0x65, 0x73, 0x74, 0x1a, 0x16, 0x2e, 0x67, 0x6f, 0x6f, 0x67, 0x6c, 0x65, 0x2e, 0x70, 0x72, 0x6f, - 0x74, 0x6f, 0x62, 0x75, 0x66, 0x2e, 0x45, 0x6d, 0x70, 0x74, 0x79, 0x12, 0x75, 0x0a, 0x1f, 0x55, - 0x70, 0x64, 0x61, 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x43, 0x6f, - 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x49, 0x6e, 0x66, 0x6f, 0x12, 0x3a, - 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, - 0x2e, 0x76, 0x31, 0x2e, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, - 0x61, 0x6e, 0x74, 0x43, 0x6f, 0x6e, 0x76, 0x65, 0x72, 0x73, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x49, - 0x6e, 0x66, 0x6f, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x16, 0x2e, 0x67, 0x6f, 0x6f, - 0x67, 0x6c, 0x65, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x62, 0x75, 0x66, 0x2e, 0x45, 0x6d, 0x70, - 0x74, 0x79, 0x12, 0x6a, 0x0a, 0x0f, 0x49, 0x73, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x45, 0x6e, - 0x61, 0x62, 0x6c, 0x65, 0x64, 0x12, 0x2a, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, - 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x49, 0x73, 0x41, 0x73, 0x73, - 0x69, 0x73, 0x74, 0x45, 0x6e, 0x61, 0x62, 0x6c, 0x65, 0x64, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, - 0x74, 0x1a, 0x2b, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, - 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x49, 0x73, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x45, - 0x6e, 0x61, 0x62, 0x6c, 0x65, 0x64, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x7f, - 0x0a, 0x16, 0x53, 0x65, 0x61, 0x72, 0x63, 0x68, 0x55, 0x6e, 0x69, 0x66, 0x69, 0x65, 0x64, 0x52, - 0x65, 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x73, 0x12, 0x31, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, - 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x53, 0x65, - 0x61, 0x72, 0x63, 0x68, 0x55, 0x6e, 0x69, 0x66, 0x69, 0x65, 0x64, 0x52, 0x65, 0x73, 0x6f, 0x75, - 0x72, 0x63, 0x65, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x32, 0x2e, 0x74, 0x65, - 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, - 0x2e, 0x53, 0x65, 0x61, 0x72, 0x63, 0x68, 0x55, 0x6e, 0x69, 0x66, 0x69, 0x65, 0x64, 0x52, 0x65, - 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x32, - 0x99, 0x01, 0x0a, 0x16, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x45, 0x6d, 0x62, 0x65, 0x64, 0x64, - 0x69, 0x6e, 0x67, 0x53, 0x65, 0x72, 0x76, 0x69, 0x63, 0x65, 0x12, 0x7f, 0x0a, 0x16, 0x47, 0x65, - 0x74, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x45, 0x6d, 0x62, 0x65, 0x64, 0x64, - 0x69, 0x6e, 0x67, 0x73, 0x12, 0x31, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, - 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x47, 0x65, 0x74, 0x41, 0x73, 0x73, - 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x45, 0x6d, 0x62, 0x65, 0x64, 0x64, 0x69, 0x6e, 0x67, 0x73, - 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x32, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, - 0x72, 0x74, 0x2e, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2e, 0x76, 0x31, 0x2e, 0x47, 0x65, 0x74, - 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x61, 0x6e, 0x74, 0x45, 0x6d, 0x62, 0x65, 0x64, 0x64, 0x69, - 0x6e, 0x67, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x42, 0x45, 0x5a, 0x43, 0x67, - 0x69, 0x74, 0x68, 0x75, 0x62, 0x2e, 0x63, 0x6f, 0x6d, 0x2f, 0x67, 0x72, 0x61, 0x76, 0x69, 0x74, - 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x61, 0x6c, 0x2f, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, - 0x2f, 0x61, 0x70, 0x69, 0x2f, 0x67, 0x65, 0x6e, 0x2f, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x2f, 0x67, - 0x6f, 0x2f, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, 0x2f, 0x76, 0x31, 0x3b, 0x61, 0x73, 0x73, 0x69, - 0x73, 0x74, 0x62, 0x06, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x33, -} - -var ( - file_teleport_assist_v1_assist_proto_rawDescOnce sync.Once - file_teleport_assist_v1_assist_proto_rawDescData = file_teleport_assist_v1_assist_proto_rawDesc -) - -func file_teleport_assist_v1_assist_proto_rawDescGZIP() []byte { - file_teleport_assist_v1_assist_proto_rawDescOnce.Do(func() { - file_teleport_assist_v1_assist_proto_rawDescData = protoimpl.X.CompressGZIP(file_teleport_assist_v1_assist_proto_rawDescData) - }) - return file_teleport_assist_v1_assist_proto_rawDescData -} - -var file_teleport_assist_v1_assist_proto_msgTypes = make([]protoimpl.MessageInfo, 18) -var file_teleport_assist_v1_assist_proto_goTypes = []interface{}{ - (*GetAssistantMessagesRequest)(nil), // 0: teleport.assist.v1.GetAssistantMessagesRequest - (*AssistantMessage)(nil), // 1: teleport.assist.v1.AssistantMessage - (*CreateAssistantMessageRequest)(nil), // 2: teleport.assist.v1.CreateAssistantMessageRequest - (*GetAssistantMessagesResponse)(nil), // 3: teleport.assist.v1.GetAssistantMessagesResponse - (*GetAssistantConversationsRequest)(nil), // 4: teleport.assist.v1.GetAssistantConversationsRequest - (*ConversationInfo)(nil), // 5: teleport.assist.v1.ConversationInfo - (*GetAssistantConversationsResponse)(nil), // 6: teleport.assist.v1.GetAssistantConversationsResponse - (*CreateAssistantConversationRequest)(nil), // 7: teleport.assist.v1.CreateAssistantConversationRequest - (*CreateAssistantConversationResponse)(nil), // 8: teleport.assist.v1.CreateAssistantConversationResponse - (*UpdateAssistantConversationInfoRequest)(nil), // 9: teleport.assist.v1.UpdateAssistantConversationInfoRequest - (*IsAssistEnabledRequest)(nil), // 10: teleport.assist.v1.IsAssistEnabledRequest - (*IsAssistEnabledResponse)(nil), // 11: teleport.assist.v1.IsAssistEnabledResponse - (*DeleteAssistantConversationRequest)(nil), // 12: teleport.assist.v1.DeleteAssistantConversationRequest - (*GetAssistantEmbeddingsRequest)(nil), // 13: teleport.assist.v1.GetAssistantEmbeddingsRequest - (*EmbeddedDocument)(nil), // 14: teleport.assist.v1.EmbeddedDocument - (*GetAssistantEmbeddingsResponse)(nil), // 15: teleport.assist.v1.GetAssistantEmbeddingsResponse - (*SearchUnifiedResourcesRequest)(nil), // 16: teleport.assist.v1.SearchUnifiedResourcesRequest - (*SearchUnifiedResourcesResponse)(nil), // 17: teleport.assist.v1.SearchUnifiedResourcesResponse - (*timestamppb.Timestamp)(nil), // 18: google.protobuf.Timestamp - (*proto.PaginatedResource)(nil), // 19: proto.PaginatedResource - (*emptypb.Empty)(nil), // 20: google.protobuf.Empty -} -var file_teleport_assist_v1_assist_proto_depIdxs = []int32{ - 18, // 0: teleport.assist.v1.AssistantMessage.created_time:type_name -> google.protobuf.Timestamp - 1, // 1: teleport.assist.v1.CreateAssistantMessageRequest.message:type_name -> teleport.assist.v1.AssistantMessage - 1, // 2: teleport.assist.v1.GetAssistantMessagesResponse.messages:type_name -> teleport.assist.v1.AssistantMessage - 18, // 3: teleport.assist.v1.ConversationInfo.created_time:type_name -> google.protobuf.Timestamp - 5, // 4: teleport.assist.v1.GetAssistantConversationsResponse.conversations:type_name -> teleport.assist.v1.ConversationInfo - 18, // 5: teleport.assist.v1.CreateAssistantConversationRequest.created_time:type_name -> google.protobuf.Timestamp - 14, // 6: teleport.assist.v1.GetAssistantEmbeddingsResponse.embeddings:type_name -> teleport.assist.v1.EmbeddedDocument - 19, // 7: teleport.assist.v1.SearchUnifiedResourcesResponse.resources:type_name -> proto.PaginatedResource - 7, // 8: teleport.assist.v1.AssistService.CreateAssistantConversation:input_type -> teleport.assist.v1.CreateAssistantConversationRequest - 4, // 9: teleport.assist.v1.AssistService.GetAssistantConversations:input_type -> teleport.assist.v1.GetAssistantConversationsRequest - 12, // 10: teleport.assist.v1.AssistService.DeleteAssistantConversation:input_type -> teleport.assist.v1.DeleteAssistantConversationRequest - 0, // 11: teleport.assist.v1.AssistService.GetAssistantMessages:input_type -> teleport.assist.v1.GetAssistantMessagesRequest - 2, // 12: teleport.assist.v1.AssistService.CreateAssistantMessage:input_type -> teleport.assist.v1.CreateAssistantMessageRequest - 9, // 13: teleport.assist.v1.AssistService.UpdateAssistantConversationInfo:input_type -> teleport.assist.v1.UpdateAssistantConversationInfoRequest - 10, // 14: teleport.assist.v1.AssistService.IsAssistEnabled:input_type -> teleport.assist.v1.IsAssistEnabledRequest - 16, // 15: teleport.assist.v1.AssistService.SearchUnifiedResources:input_type -> teleport.assist.v1.SearchUnifiedResourcesRequest - 13, // 16: teleport.assist.v1.AssistEmbeddingService.GetAssistantEmbeddings:input_type -> teleport.assist.v1.GetAssistantEmbeddingsRequest - 8, // 17: teleport.assist.v1.AssistService.CreateAssistantConversation:output_type -> teleport.assist.v1.CreateAssistantConversationResponse - 6, // 18: teleport.assist.v1.AssistService.GetAssistantConversations:output_type -> teleport.assist.v1.GetAssistantConversationsResponse - 20, // 19: teleport.assist.v1.AssistService.DeleteAssistantConversation:output_type -> google.protobuf.Empty - 3, // 20: teleport.assist.v1.AssistService.GetAssistantMessages:output_type -> teleport.assist.v1.GetAssistantMessagesResponse - 20, // 21: teleport.assist.v1.AssistService.CreateAssistantMessage:output_type -> google.protobuf.Empty - 20, // 22: teleport.assist.v1.AssistService.UpdateAssistantConversationInfo:output_type -> google.protobuf.Empty - 11, // 23: teleport.assist.v1.AssistService.IsAssistEnabled:output_type -> teleport.assist.v1.IsAssistEnabledResponse - 17, // 24: teleport.assist.v1.AssistService.SearchUnifiedResources:output_type -> teleport.assist.v1.SearchUnifiedResourcesResponse - 15, // 25: teleport.assist.v1.AssistEmbeddingService.GetAssistantEmbeddings:output_type -> teleport.assist.v1.GetAssistantEmbeddingsResponse - 17, // [17:26] is the sub-list for method output_type - 8, // [8:17] is the sub-list for method input_type - 8, // [8:8] is the sub-list for extension type_name - 8, // [8:8] is the sub-list for extension extendee - 0, // [0:8] is the sub-list for field type_name -} - -func init() { file_teleport_assist_v1_assist_proto_init() } -func file_teleport_assist_v1_assist_proto_init() { - if File_teleport_assist_v1_assist_proto != nil { - return - } - if !protoimpl.UnsafeEnabled { - file_teleport_assist_v1_assist_proto_msgTypes[0].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*GetAssistantMessagesRequest); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[1].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*AssistantMessage); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[2].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*CreateAssistantMessageRequest); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[3].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*GetAssistantMessagesResponse); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[4].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*GetAssistantConversationsRequest); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[5].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*ConversationInfo); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[6].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*GetAssistantConversationsResponse); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[7].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*CreateAssistantConversationRequest); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[8].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*CreateAssistantConversationResponse); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[9].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*UpdateAssistantConversationInfoRequest); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[10].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*IsAssistEnabledRequest); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[11].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*IsAssistEnabledResponse); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[12].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*DeleteAssistantConversationRequest); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[13].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*GetAssistantEmbeddingsRequest); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[14].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*EmbeddedDocument); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[15].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*GetAssistantEmbeddingsResponse); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[16].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*SearchUnifiedResourcesRequest); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - file_teleport_assist_v1_assist_proto_msgTypes[17].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*SearchUnifiedResourcesResponse); i { - case 0: - return &v.state - case 1: - return &v.sizeCache - case 2: - return &v.unknownFields - default: - return nil - } - } - } - type x struct{} - out := protoimpl.TypeBuilder{ - File: protoimpl.DescBuilder{ - GoPackagePath: reflect.TypeOf(x{}).PkgPath(), - RawDescriptor: file_teleport_assist_v1_assist_proto_rawDesc, - NumEnums: 0, - NumMessages: 18, - NumExtensions: 0, - NumServices: 2, - }, - GoTypes: file_teleport_assist_v1_assist_proto_goTypes, - DependencyIndexes: file_teleport_assist_v1_assist_proto_depIdxs, - MessageInfos: file_teleport_assist_v1_assist_proto_msgTypes, - }.Build() - File_teleport_assist_v1_assist_proto = out.File - file_teleport_assist_v1_assist_proto_rawDesc = nil - file_teleport_assist_v1_assist_proto_goTypes = nil - file_teleport_assist_v1_assist_proto_depIdxs = nil -} diff --git a/api/gen/proto/go/assist/v1/assist_grpc.pb.go b/api/gen/proto/go/assist/v1/assist_grpc.pb.go deleted file mode 100644 index c9b63831db0..00000000000 --- a/api/gen/proto/go/assist/v1/assist_grpc.pb.go +++ /dev/null @@ -1,509 +0,0 @@ -// Copyright 2023 Gravitational, Inc -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -// Code generated by protoc-gen-go-grpc. DO NOT EDIT. -// versions: -// - protoc-gen-go-grpc v1.4.0 -// - protoc (unknown) -// source: teleport/assist/v1/assist.proto - -package assist - -import ( - context "context" - grpc "google.golang.org/grpc" - codes "google.golang.org/grpc/codes" - status "google.golang.org/grpc/status" - emptypb "google.golang.org/protobuf/types/known/emptypb" -) - -// This is a compile-time assertion to ensure that this generated file -// is compatible with the grpc package it is being compiled against. -// Requires gRPC-Go v1.62.0 or later. -const _ = grpc.SupportPackageIsVersion8 - -const ( - AssistService_CreateAssistantConversation_FullMethodName = "/teleport.assist.v1.AssistService/CreateAssistantConversation" - AssistService_GetAssistantConversations_FullMethodName = "/teleport.assist.v1.AssistService/GetAssistantConversations" - AssistService_DeleteAssistantConversation_FullMethodName = "/teleport.assist.v1.AssistService/DeleteAssistantConversation" - AssistService_GetAssistantMessages_FullMethodName = "/teleport.assist.v1.AssistService/GetAssistantMessages" - AssistService_CreateAssistantMessage_FullMethodName = "/teleport.assist.v1.AssistService/CreateAssistantMessage" - AssistService_UpdateAssistantConversationInfo_FullMethodName = "/teleport.assist.v1.AssistService/UpdateAssistantConversationInfo" - AssistService_IsAssistEnabled_FullMethodName = "/teleport.assist.v1.AssistService/IsAssistEnabled" - AssistService_SearchUnifiedResources_FullMethodName = "/teleport.assist.v1.AssistService/SearchUnifiedResources" -) - -// AssistServiceClient is the client API for AssistService service. -// -// For semantics around ctx use and closing/ending streaming RPCs, please refer to https://pkg.go.dev/google.golang.org/grpc/?tab=doc#ClientConn.NewStream. -// -// AssistService is a service that provides an ability to communicate with the Teleport Assist. -type AssistServiceClient interface { - // CreateNewConversation creates a new conversation and returns the UUID of it. - CreateAssistantConversation(ctx context.Context, in *CreateAssistantConversationRequest, opts ...grpc.CallOption) (*CreateAssistantConversationResponse, error) - // GetAssistantConversations returns all conversations for the connected user. - GetAssistantConversations(ctx context.Context, in *GetAssistantConversationsRequest, opts ...grpc.CallOption) (*GetAssistantConversationsResponse, error) - // DeleteAssistantConversation deletes the conversation and all messages associated with it. - DeleteAssistantConversation(ctx context.Context, in *DeleteAssistantConversationRequest, opts ...grpc.CallOption) (*emptypb.Empty, error) - // GetAssistantMessages returns all messages associated with the given conversation ID. - GetAssistantMessages(ctx context.Context, in *GetAssistantMessagesRequest, opts ...grpc.CallOption) (*GetAssistantMessagesResponse, error) - // CreateAssistantMessage creates a new message in the given conversation. - CreateAssistantMessage(ctx context.Context, in *CreateAssistantMessageRequest, opts ...grpc.CallOption) (*emptypb.Empty, error) - // UpdateAssistantConversationInfo updates the conversation info. - UpdateAssistantConversationInfo(ctx context.Context, in *UpdateAssistantConversationInfoRequest, opts ...grpc.CallOption) (*emptypb.Empty, error) - // IsAssistEnabled returns true if the assist is enabled or not on the auth level. - IsAssistEnabled(ctx context.Context, in *IsAssistEnabledRequest, opts ...grpc.CallOption) (*IsAssistEnabledResponse, error) - // SearchUnifiedResources returns a similarity-ordered list of resources from the unified resource cache. - SearchUnifiedResources(ctx context.Context, in *SearchUnifiedResourcesRequest, opts ...grpc.CallOption) (*SearchUnifiedResourcesResponse, error) -} - -type assistServiceClient struct { - cc grpc.ClientConnInterface -} - -func NewAssistServiceClient(cc grpc.ClientConnInterface) AssistServiceClient { - return &assistServiceClient{cc} -} - -func (c *assistServiceClient) CreateAssistantConversation(ctx context.Context, in *CreateAssistantConversationRequest, opts ...grpc.CallOption) (*CreateAssistantConversationResponse, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(CreateAssistantConversationResponse) - err := c.cc.Invoke(ctx, AssistService_CreateAssistantConversation_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -func (c *assistServiceClient) GetAssistantConversations(ctx context.Context, in *GetAssistantConversationsRequest, opts ...grpc.CallOption) (*GetAssistantConversationsResponse, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(GetAssistantConversationsResponse) - err := c.cc.Invoke(ctx, AssistService_GetAssistantConversations_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -func (c *assistServiceClient) DeleteAssistantConversation(ctx context.Context, in *DeleteAssistantConversationRequest, opts ...grpc.CallOption) (*emptypb.Empty, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(emptypb.Empty) - err := c.cc.Invoke(ctx, AssistService_DeleteAssistantConversation_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -func (c *assistServiceClient) GetAssistantMessages(ctx context.Context, in *GetAssistantMessagesRequest, opts ...grpc.CallOption) (*GetAssistantMessagesResponse, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(GetAssistantMessagesResponse) - err := c.cc.Invoke(ctx, AssistService_GetAssistantMessages_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -func (c *assistServiceClient) CreateAssistantMessage(ctx context.Context, in *CreateAssistantMessageRequest, opts ...grpc.CallOption) (*emptypb.Empty, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(emptypb.Empty) - err := c.cc.Invoke(ctx, AssistService_CreateAssistantMessage_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -func (c *assistServiceClient) UpdateAssistantConversationInfo(ctx context.Context, in *UpdateAssistantConversationInfoRequest, opts ...grpc.CallOption) (*emptypb.Empty, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(emptypb.Empty) - err := c.cc.Invoke(ctx, AssistService_UpdateAssistantConversationInfo_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -func (c *assistServiceClient) IsAssistEnabled(ctx context.Context, in *IsAssistEnabledRequest, opts ...grpc.CallOption) (*IsAssistEnabledResponse, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(IsAssistEnabledResponse) - err := c.cc.Invoke(ctx, AssistService_IsAssistEnabled_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -func (c *assistServiceClient) SearchUnifiedResources(ctx context.Context, in *SearchUnifiedResourcesRequest, opts ...grpc.CallOption) (*SearchUnifiedResourcesResponse, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(SearchUnifiedResourcesResponse) - err := c.cc.Invoke(ctx, AssistService_SearchUnifiedResources_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -// AssistServiceServer is the server API for AssistService service. -// All implementations must embed UnimplementedAssistServiceServer -// for forward compatibility -// -// AssistService is a service that provides an ability to communicate with the Teleport Assist. -type AssistServiceServer interface { - // CreateNewConversation creates a new conversation and returns the UUID of it. - CreateAssistantConversation(context.Context, *CreateAssistantConversationRequest) (*CreateAssistantConversationResponse, error) - // GetAssistantConversations returns all conversations for the connected user. - GetAssistantConversations(context.Context, *GetAssistantConversationsRequest) (*GetAssistantConversationsResponse, error) - // DeleteAssistantConversation deletes the conversation and all messages associated with it. - DeleteAssistantConversation(context.Context, *DeleteAssistantConversationRequest) (*emptypb.Empty, error) - // GetAssistantMessages returns all messages associated with the given conversation ID. - GetAssistantMessages(context.Context, *GetAssistantMessagesRequest) (*GetAssistantMessagesResponse, error) - // CreateAssistantMessage creates a new message in the given conversation. - CreateAssistantMessage(context.Context, *CreateAssistantMessageRequest) (*emptypb.Empty, error) - // UpdateAssistantConversationInfo updates the conversation info. - UpdateAssistantConversationInfo(context.Context, *UpdateAssistantConversationInfoRequest) (*emptypb.Empty, error) - // IsAssistEnabled returns true if the assist is enabled or not on the auth level. - IsAssistEnabled(context.Context, *IsAssistEnabledRequest) (*IsAssistEnabledResponse, error) - // SearchUnifiedResources returns a similarity-ordered list of resources from the unified resource cache. - SearchUnifiedResources(context.Context, *SearchUnifiedResourcesRequest) (*SearchUnifiedResourcesResponse, error) - mustEmbedUnimplementedAssistServiceServer() -} - -// UnimplementedAssistServiceServer must be embedded to have forward compatible implementations. -type UnimplementedAssistServiceServer struct { -} - -func (UnimplementedAssistServiceServer) CreateAssistantConversation(context.Context, *CreateAssistantConversationRequest) (*CreateAssistantConversationResponse, error) { - return nil, status.Errorf(codes.Unimplemented, "method CreateAssistantConversation not implemented") -} -func (UnimplementedAssistServiceServer) GetAssistantConversations(context.Context, *GetAssistantConversationsRequest) (*GetAssistantConversationsResponse, error) { - return nil, status.Errorf(codes.Unimplemented, "method GetAssistantConversations not implemented") -} -func (UnimplementedAssistServiceServer) DeleteAssistantConversation(context.Context, *DeleteAssistantConversationRequest) (*emptypb.Empty, error) { - return nil, status.Errorf(codes.Unimplemented, "method DeleteAssistantConversation not implemented") -} -func (UnimplementedAssistServiceServer) GetAssistantMessages(context.Context, *GetAssistantMessagesRequest) (*GetAssistantMessagesResponse, error) { - return nil, status.Errorf(codes.Unimplemented, "method GetAssistantMessages not implemented") -} -func (UnimplementedAssistServiceServer) CreateAssistantMessage(context.Context, *CreateAssistantMessageRequest) (*emptypb.Empty, error) { - return nil, status.Errorf(codes.Unimplemented, "method CreateAssistantMessage not implemented") -} -func (UnimplementedAssistServiceServer) UpdateAssistantConversationInfo(context.Context, *UpdateAssistantConversationInfoRequest) (*emptypb.Empty, error) { - return nil, status.Errorf(codes.Unimplemented, "method UpdateAssistantConversationInfo not implemented") -} -func (UnimplementedAssistServiceServer) IsAssistEnabled(context.Context, *IsAssistEnabledRequest) (*IsAssistEnabledResponse, error) { - return nil, status.Errorf(codes.Unimplemented, "method IsAssistEnabled not implemented") -} -func (UnimplementedAssistServiceServer) SearchUnifiedResources(context.Context, *SearchUnifiedResourcesRequest) (*SearchUnifiedResourcesResponse, error) { - return nil, status.Errorf(codes.Unimplemented, "method SearchUnifiedResources not implemented") -} -func (UnimplementedAssistServiceServer) mustEmbedUnimplementedAssistServiceServer() {} - -// UnsafeAssistServiceServer may be embedded to opt out of forward compatibility for this service. -// Use of this interface is not recommended, as added methods to AssistServiceServer will -// result in compilation errors. -type UnsafeAssistServiceServer interface { - mustEmbedUnimplementedAssistServiceServer() -} - -func RegisterAssistServiceServer(s grpc.ServiceRegistrar, srv AssistServiceServer) { - s.RegisterService(&AssistService_ServiceDesc, srv) -} - -func _AssistService_CreateAssistantConversation_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(CreateAssistantConversationRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AssistServiceServer).CreateAssistantConversation(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AssistService_CreateAssistantConversation_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AssistServiceServer).CreateAssistantConversation(ctx, req.(*CreateAssistantConversationRequest)) - } - return interceptor(ctx, in, info, handler) -} - -func _AssistService_GetAssistantConversations_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(GetAssistantConversationsRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AssistServiceServer).GetAssistantConversations(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AssistService_GetAssistantConversations_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AssistServiceServer).GetAssistantConversations(ctx, req.(*GetAssistantConversationsRequest)) - } - return interceptor(ctx, in, info, handler) -} - -func _AssistService_DeleteAssistantConversation_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(DeleteAssistantConversationRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AssistServiceServer).DeleteAssistantConversation(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AssistService_DeleteAssistantConversation_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AssistServiceServer).DeleteAssistantConversation(ctx, req.(*DeleteAssistantConversationRequest)) - } - return interceptor(ctx, in, info, handler) -} - -func _AssistService_GetAssistantMessages_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(GetAssistantMessagesRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AssistServiceServer).GetAssistantMessages(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AssistService_GetAssistantMessages_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AssistServiceServer).GetAssistantMessages(ctx, req.(*GetAssistantMessagesRequest)) - } - return interceptor(ctx, in, info, handler) -} - -func _AssistService_CreateAssistantMessage_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(CreateAssistantMessageRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AssistServiceServer).CreateAssistantMessage(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AssistService_CreateAssistantMessage_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AssistServiceServer).CreateAssistantMessage(ctx, req.(*CreateAssistantMessageRequest)) - } - return interceptor(ctx, in, info, handler) -} - -func _AssistService_UpdateAssistantConversationInfo_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(UpdateAssistantConversationInfoRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AssistServiceServer).UpdateAssistantConversationInfo(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AssistService_UpdateAssistantConversationInfo_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AssistServiceServer).UpdateAssistantConversationInfo(ctx, req.(*UpdateAssistantConversationInfoRequest)) - } - return interceptor(ctx, in, info, handler) -} - -func _AssistService_IsAssistEnabled_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(IsAssistEnabledRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AssistServiceServer).IsAssistEnabled(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AssistService_IsAssistEnabled_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AssistServiceServer).IsAssistEnabled(ctx, req.(*IsAssistEnabledRequest)) - } - return interceptor(ctx, in, info, handler) -} - -func _AssistService_SearchUnifiedResources_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(SearchUnifiedResourcesRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AssistServiceServer).SearchUnifiedResources(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AssistService_SearchUnifiedResources_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AssistServiceServer).SearchUnifiedResources(ctx, req.(*SearchUnifiedResourcesRequest)) - } - return interceptor(ctx, in, info, handler) -} - -// AssistService_ServiceDesc is the grpc.ServiceDesc for AssistService service. -// It's only intended for direct use with grpc.RegisterService, -// and not to be introspected or modified (even as a copy) -var AssistService_ServiceDesc = grpc.ServiceDesc{ - ServiceName: "teleport.assist.v1.AssistService", - HandlerType: (*AssistServiceServer)(nil), - Methods: []grpc.MethodDesc{ - { - MethodName: "CreateAssistantConversation", - Handler: _AssistService_CreateAssistantConversation_Handler, - }, - { - MethodName: "GetAssistantConversations", - Handler: _AssistService_GetAssistantConversations_Handler, - }, - { - MethodName: "DeleteAssistantConversation", - Handler: _AssistService_DeleteAssistantConversation_Handler, - }, - { - MethodName: "GetAssistantMessages", - Handler: _AssistService_GetAssistantMessages_Handler, - }, - { - MethodName: "CreateAssistantMessage", - Handler: _AssistService_CreateAssistantMessage_Handler, - }, - { - MethodName: "UpdateAssistantConversationInfo", - Handler: _AssistService_UpdateAssistantConversationInfo_Handler, - }, - { - MethodName: "IsAssistEnabled", - Handler: _AssistService_IsAssistEnabled_Handler, - }, - { - MethodName: "SearchUnifiedResources", - Handler: _AssistService_SearchUnifiedResources_Handler, - }, - }, - Streams: []grpc.StreamDesc{}, - Metadata: "teleport/assist/v1/assist.proto", -} - -const ( - AssistEmbeddingService_GetAssistantEmbeddings_FullMethodName = "/teleport.assist.v1.AssistEmbeddingService/GetAssistantEmbeddings" -) - -// AssistEmbeddingServiceClient is the client API for AssistEmbeddingService service. -// -// For semantics around ctx use and closing/ending streaming RPCs, please refer to https://pkg.go.dev/google.golang.org/grpc/?tab=doc#ClientConn.NewStream. -// -// AssistEmbeddingService is a service that provides an ability to communicate with the Assist Embedding service. -type AssistEmbeddingServiceClient interface { - // AssistantGetEmbeddings returns the embeddings for the given query. - GetAssistantEmbeddings(ctx context.Context, in *GetAssistantEmbeddingsRequest, opts ...grpc.CallOption) (*GetAssistantEmbeddingsResponse, error) -} - -type assistEmbeddingServiceClient struct { - cc grpc.ClientConnInterface -} - -func NewAssistEmbeddingServiceClient(cc grpc.ClientConnInterface) AssistEmbeddingServiceClient { - return &assistEmbeddingServiceClient{cc} -} - -func (c *assistEmbeddingServiceClient) GetAssistantEmbeddings(ctx context.Context, in *GetAssistantEmbeddingsRequest, opts ...grpc.CallOption) (*GetAssistantEmbeddingsResponse, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(GetAssistantEmbeddingsResponse) - err := c.cc.Invoke(ctx, AssistEmbeddingService_GetAssistantEmbeddings_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -// AssistEmbeddingServiceServer is the server API for AssistEmbeddingService service. -// All implementations must embed UnimplementedAssistEmbeddingServiceServer -// for forward compatibility -// -// AssistEmbeddingService is a service that provides an ability to communicate with the Assist Embedding service. -type AssistEmbeddingServiceServer interface { - // AssistantGetEmbeddings returns the embeddings for the given query. - GetAssistantEmbeddings(context.Context, *GetAssistantEmbeddingsRequest) (*GetAssistantEmbeddingsResponse, error) - mustEmbedUnimplementedAssistEmbeddingServiceServer() -} - -// UnimplementedAssistEmbeddingServiceServer must be embedded to have forward compatible implementations. -type UnimplementedAssistEmbeddingServiceServer struct { -} - -func (UnimplementedAssistEmbeddingServiceServer) GetAssistantEmbeddings(context.Context, *GetAssistantEmbeddingsRequest) (*GetAssistantEmbeddingsResponse, error) { - return nil, status.Errorf(codes.Unimplemented, "method GetAssistantEmbeddings not implemented") -} -func (UnimplementedAssistEmbeddingServiceServer) mustEmbedUnimplementedAssistEmbeddingServiceServer() { -} - -// UnsafeAssistEmbeddingServiceServer may be embedded to opt out of forward compatibility for this service. -// Use of this interface is not recommended, as added methods to AssistEmbeddingServiceServer will -// result in compilation errors. -type UnsafeAssistEmbeddingServiceServer interface { - mustEmbedUnimplementedAssistEmbeddingServiceServer() -} - -func RegisterAssistEmbeddingServiceServer(s grpc.ServiceRegistrar, srv AssistEmbeddingServiceServer) { - s.RegisterService(&AssistEmbeddingService_ServiceDesc, srv) -} - -func _AssistEmbeddingService_GetAssistantEmbeddings_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(GetAssistantEmbeddingsRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AssistEmbeddingServiceServer).GetAssistantEmbeddings(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AssistEmbeddingService_GetAssistantEmbeddings_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AssistEmbeddingServiceServer).GetAssistantEmbeddings(ctx, req.(*GetAssistantEmbeddingsRequest)) - } - return interceptor(ctx, in, info, handler) -} - -// AssistEmbeddingService_ServiceDesc is the grpc.ServiceDesc for AssistEmbeddingService service. -// It's only intended for direct use with grpc.RegisterService, -// and not to be introspected or modified (even as a copy) -var AssistEmbeddingService_ServiceDesc = grpc.ServiceDesc{ - ServiceName: "teleport.assist.v1.AssistEmbeddingService", - HandlerType: (*AssistEmbeddingServiceServer)(nil), - Methods: []grpc.MethodDesc{ - { - MethodName: "GetAssistantEmbeddings", - Handler: _AssistEmbeddingService_GetAssistantEmbeddings_Handler, - }, - }, - Streams: []grpc.StreamDesc{}, - Metadata: "teleport/assist/v1/assist.proto", -} diff --git a/api/gen/proto/go/userpreferences/v1/userpreferences.pb.go b/api/gen/proto/go/userpreferences/v1/userpreferences.pb.go index f5614530c06..5732048033c 100644 --- a/api/gen/proto/go/userpreferences/v1/userpreferences.pb.go +++ b/api/gen/proto/go/userpreferences/v1/userpreferences.pb.go @@ -41,8 +41,6 @@ type UserPreferences struct { sizeCache protoimpl.SizeCache unknownFields protoimpl.UnknownFields - // assist is the preferences for the Teleport Assist. - Assist *AssistUserPreferences `protobuf:"bytes,1,opt,name=assist,proto3" json:"assist,omitempty"` // theme is the theme of the frontend. Theme Theme `protobuf:"varint,2,opt,name=theme,proto3,enum=teleport.userpreferences.v1.Theme" json:"theme,omitempty"` // onboard is the preferences from the onboarding questionnaire. @@ -87,13 +85,6 @@ func (*UserPreferences) Descriptor() ([]byte, []int) { return file_teleport_userpreferences_v1_userpreferences_proto_rawDescGZIP(), []int{0} } -func (x *UserPreferences) GetAssist() *AssistUserPreferences { - if x != nil { - return x.Assist - } - return nil -} - func (x *UserPreferences) GetTheme() Theme { if x != nil { return x.Theme @@ -278,98 +269,91 @@ var file_teleport_userpreferences_v1_userpreferences_proto_rawDesc = []byte{ 0x66, 0x2f, 0x65, 0x6d, 0x70, 0x74, 0x79, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x1a, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2f, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2f, 0x76, 0x31, 0x2f, 0x61, 0x63, 0x63, 0x65, 0x73, - 0x73, 0x5f, 0x67, 0x72, 0x61, 0x70, 0x68, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x1a, 0x28, 0x74, + 0x73, 0x5f, 0x67, 0x72, 0x61, 0x70, 0x68, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x1a, 0x35, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2f, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, - 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2f, 0x76, 0x31, 0x2f, 0x61, 0x73, 0x73, 0x69, 0x73, - 0x74, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x1a, 0x35, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, - 0x74, 0x2f, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, - 0x73, 0x2f, 0x76, 0x31, 0x2f, 0x63, 0x6c, 0x75, 0x73, 0x74, 0x65, 0x72, 0x5f, 0x70, 0x72, 0x65, - 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x1a, 0x29, - 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2f, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, - 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2f, 0x76, 0x31, 0x2f, 0x6f, 0x6e, 0x62, 0x6f, - 0x61, 0x72, 0x64, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x1a, 0x27, 0x74, 0x65, 0x6c, 0x65, 0x70, - 0x6f, 0x72, 0x74, 0x2f, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, - 0x63, 0x65, 0x73, 0x2f, 0x76, 0x31, 0x2f, 0x74, 0x68, 0x65, 0x6d, 0x65, 0x2e, 0x70, 0x72, 0x6f, - 0x74, 0x6f, 0x1a, 0x3e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2f, 0x75, 0x73, 0x65, - 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2f, 0x76, 0x31, 0x2f, - 0x75, 0x6e, 0x69, 0x66, 0x69, 0x65, 0x64, 0x5f, 0x72, 0x65, 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, - 0x5f, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x70, 0x72, 0x6f, - 0x74, 0x6f, 0x22, 0xa3, 0x04, 0x0a, 0x0f, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, - 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x12, 0x4a, 0x0a, 0x06, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, - 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x32, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, - 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, - 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x41, 0x73, 0x73, 0x69, 0x73, 0x74, 0x55, 0x73, 0x65, 0x72, 0x50, - 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x06, 0x61, 0x73, 0x73, 0x69, - 0x73, 0x74, 0x12, 0x38, 0x0a, 0x05, 0x74, 0x68, 0x65, 0x6d, 0x65, 0x18, 0x02, 0x20, 0x01, 0x28, - 0x0e, 0x32, 0x22, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, - 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, - 0x54, 0x68, 0x65, 0x6d, 0x65, 0x52, 0x05, 0x74, 0x68, 0x65, 0x6d, 0x65, 0x12, 0x4d, 0x0a, 0x07, - 0x6f, 0x6e, 0x62, 0x6f, 0x61, 0x72, 0x64, 0x18, 0x03, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x33, 0x2e, - 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, - 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x4f, 0x6e, 0x62, 0x6f, - 0x61, 0x72, 0x64, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, - 0x65, 0x73, 0x52, 0x07, 0x6f, 0x6e, 0x62, 0x6f, 0x61, 0x72, 0x64, 0x12, 0x64, 0x0a, 0x13, 0x63, - 0x6c, 0x75, 0x73, 0x74, 0x65, 0x72, 0x5f, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, - 0x65, 0x73, 0x18, 0x04, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x33, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, - 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, - 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x43, 0x6c, 0x75, 0x73, 0x74, 0x65, 0x72, 0x55, 0x73, - 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x12, 0x63, - 0x6c, 0x75, 0x73, 0x74, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, - 0x73, 0x12, 0x79, 0x0a, 0x1c, 0x75, 0x6e, 0x69, 0x66, 0x69, 0x65, 0x64, 0x5f, 0x72, 0x65, 0x73, - 0x6f, 0x75, 0x72, 0x63, 0x65, 0x5f, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, - 0x73, 0x18, 0x05, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x37, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, + 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2f, 0x76, 0x31, 0x2f, 0x63, 0x6c, 0x75, 0x73, 0x74, + 0x65, 0x72, 0x5f, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x70, + 0x72, 0x6f, 0x74, 0x6f, 0x1a, 0x29, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2f, 0x75, + 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2f, 0x76, + 0x31, 0x2f, 0x6f, 0x6e, 0x62, 0x6f, 0x61, 0x72, 0x64, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x1a, + 0x27, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2f, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, + 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2f, 0x76, 0x31, 0x2f, 0x74, 0x68, 0x65, + 0x6d, 0x65, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x1a, 0x3e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, + 0x72, 0x74, 0x2f, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, + 0x65, 0x73, 0x2f, 0x76, 0x31, 0x2f, 0x75, 0x6e, 0x69, 0x66, 0x69, 0x65, 0x64, 0x5f, 0x72, 0x65, + 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x5f, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, + 0x65, 0x73, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x22, 0xe5, 0x03, 0x0a, 0x0f, 0x55, 0x73, 0x65, + 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x12, 0x38, 0x0a, 0x05, + 0x74, 0x68, 0x65, 0x6d, 0x65, 0x18, 0x02, 0x20, 0x01, 0x28, 0x0e, 0x32, 0x22, 0x2e, 0x74, 0x65, + 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, + 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x54, 0x68, 0x65, 0x6d, 0x65, 0x52, + 0x05, 0x74, 0x68, 0x65, 0x6d, 0x65, 0x12, 0x4d, 0x0a, 0x07, 0x6f, 0x6e, 0x62, 0x6f, 0x61, 0x72, + 0x64, 0x18, 0x03, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x33, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, - 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x55, 0x6e, 0x69, 0x66, 0x69, 0x65, 0x64, 0x52, 0x65, 0x73, - 0x6f, 0x75, 0x72, 0x63, 0x65, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, - 0x52, 0x1a, 0x75, 0x6e, 0x69, 0x66, 0x69, 0x65, 0x64, 0x52, 0x65, 0x73, 0x6f, 0x75, 0x72, 0x63, - 0x65, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x12, 0x5a, 0x0a, 0x0c, - 0x61, 0x63, 0x63, 0x65, 0x73, 0x73, 0x5f, 0x67, 0x72, 0x61, 0x70, 0x68, 0x18, 0x06, 0x20, 0x01, - 0x28, 0x0b, 0x32, 0x37, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, + 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x4f, 0x6e, 0x62, 0x6f, 0x61, 0x72, 0x64, 0x55, 0x73, 0x65, + 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x07, 0x6f, 0x6e, + 0x62, 0x6f, 0x61, 0x72, 0x64, 0x12, 0x64, 0x0a, 0x13, 0x63, 0x6c, 0x75, 0x73, 0x74, 0x65, 0x72, + 0x5f, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x18, 0x04, 0x20, 0x01, + 0x28, 0x0b, 0x32, 0x33, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, - 0x2e, 0x41, 0x63, 0x63, 0x65, 0x73, 0x73, 0x47, 0x72, 0x61, 0x70, 0x68, 0x55, 0x73, 0x65, 0x72, - 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x0b, 0x61, 0x63, 0x63, - 0x65, 0x73, 0x73, 0x47, 0x72, 0x61, 0x70, 0x68, 0x22, 0x2b, 0x0a, 0x19, 0x47, 0x65, 0x74, 0x55, - 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x65, - 0x71, 0x75, 0x65, 0x73, 0x74, 0x4a, 0x04, 0x08, 0x01, 0x10, 0x02, 0x52, 0x08, 0x75, 0x73, 0x65, - 0x72, 0x6e, 0x61, 0x6d, 0x65, 0x22, 0x6c, 0x0a, 0x1a, 0x47, 0x65, 0x74, 0x55, 0x73, 0x65, 0x72, - 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, - 0x6e, 0x73, 0x65, 0x12, 0x4e, 0x0a, 0x0b, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, - 0x65, 0x73, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x2c, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, - 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, - 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, - 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x0b, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, - 0x63, 0x65, 0x73, 0x22, 0x7e, 0x0a, 0x1c, 0x55, 0x70, 0x73, 0x65, 0x72, 0x74, 0x55, 0x73, 0x65, - 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x65, 0x71, 0x75, - 0x65, 0x73, 0x74, 0x12, 0x4e, 0x0a, 0x0b, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, - 0x65, 0x73, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x2c, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, - 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, - 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, - 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x0b, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, - 0x63, 0x65, 0x73, 0x4a, 0x04, 0x08, 0x02, 0x10, 0x03, 0x52, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, - 0x61, 0x6d, 0x65, 0x32, 0x8c, 0x02, 0x0a, 0x16, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, - 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x53, 0x65, 0x72, 0x76, 0x69, 0x63, 0x65, 0x12, 0x85, - 0x01, 0x0a, 0x12, 0x47, 0x65, 0x74, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, - 0x65, 0x6e, 0x63, 0x65, 0x73, 0x12, 0x36, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, - 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, - 0x2e, 0x76, 0x31, 0x2e, 0x47, 0x65, 0x74, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, - 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x37, 0x2e, + 0x2e, 0x43, 0x6c, 0x75, 0x73, 0x74, 0x65, 0x72, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, + 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x12, 0x63, 0x6c, 0x75, 0x73, 0x74, 0x65, 0x72, + 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x12, 0x79, 0x0a, 0x1c, 0x75, + 0x6e, 0x69, 0x66, 0x69, 0x65, 0x64, 0x5f, 0x72, 0x65, 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x5f, + 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x18, 0x05, 0x20, 0x01, 0x28, + 0x0b, 0x32, 0x37, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, + 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, + 0x55, 0x6e, 0x69, 0x66, 0x69, 0x65, 0x64, 0x52, 0x65, 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x50, + 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x1a, 0x75, 0x6e, 0x69, 0x66, + 0x69, 0x65, 0x64, 0x52, 0x65, 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x50, 0x72, 0x65, 0x66, 0x65, + 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x12, 0x5a, 0x0a, 0x0c, 0x61, 0x63, 0x63, 0x65, 0x73, 0x73, + 0x5f, 0x67, 0x72, 0x61, 0x70, 0x68, 0x18, 0x06, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x37, 0x2e, 0x74, + 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, + 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x41, 0x63, 0x63, 0x65, 0x73, + 0x73, 0x47, 0x72, 0x61, 0x70, 0x68, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, + 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x0b, 0x61, 0x63, 0x63, 0x65, 0x73, 0x73, 0x47, 0x72, 0x61, + 0x70, 0x68, 0x4a, 0x04, 0x08, 0x01, 0x10, 0x02, 0x52, 0x06, 0x61, 0x73, 0x73, 0x69, 0x73, 0x74, + 0x22, 0x2b, 0x0a, 0x19, 0x47, 0x65, 0x74, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, + 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x4a, 0x04, 0x08, + 0x01, 0x10, 0x02, 0x52, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, 0x65, 0x22, 0x6c, 0x0a, + 0x1a, 0x47, 0x65, 0x74, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, + 0x63, 0x65, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x4e, 0x0a, 0x0b, 0x70, + 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, + 0x32, 0x2c, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, + 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x55, + 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x0b, + 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x22, 0x7e, 0x0a, 0x1c, 0x55, + 0x70, 0x73, 0x65, 0x72, 0x74, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, + 0x6e, 0x63, 0x65, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x4e, 0x0a, 0x0b, 0x70, + 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, + 0x32, 0x2c, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, + 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x55, + 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x0b, + 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x4a, 0x04, 0x08, 0x02, 0x10, + 0x03, 0x52, 0x08, 0x75, 0x73, 0x65, 0x72, 0x6e, 0x61, 0x6d, 0x65, 0x32, 0x8c, 0x02, 0x0a, 0x16, + 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x53, + 0x65, 0x72, 0x76, 0x69, 0x63, 0x65, 0x12, 0x85, 0x01, 0x0a, 0x12, 0x47, 0x65, 0x74, 0x55, 0x73, + 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x12, 0x36, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x47, 0x65, 0x74, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x65, - 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x6a, 0x0a, 0x15, 0x55, 0x70, 0x73, 0x65, 0x72, 0x74, - 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x12, - 0x39, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, - 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x55, 0x70, - 0x73, 0x65, 0x72, 0x74, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, - 0x63, 0x65, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x16, 0x2e, 0x67, 0x6f, 0x6f, - 0x67, 0x6c, 0x65, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x62, 0x75, 0x66, 0x2e, 0x45, 0x6d, 0x70, - 0x74, 0x79, 0x42, 0x59, 0x5a, 0x57, 0x67, 0x69, 0x74, 0x68, 0x75, 0x62, 0x2e, 0x63, 0x6f, 0x6d, - 0x2f, 0x67, 0x72, 0x61, 0x76, 0x69, 0x74, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x61, 0x6c, 0x2f, 0x74, - 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2f, 0x61, 0x70, 0x69, 0x2f, 0x67, 0x65, 0x6e, 0x2f, - 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x2f, 0x67, 0x6f, 0x2f, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, - 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x2f, 0x76, 0x31, 0x3b, 0x75, 0x73, 0x65, 0x72, - 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x76, 0x31, 0x62, 0x06, 0x70, - 0x72, 0x6f, 0x74, 0x6f, 0x33, + 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x37, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, + 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, + 0x2e, 0x76, 0x31, 0x2e, 0x47, 0x65, 0x74, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, 0x65, + 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x6a, + 0x0a, 0x15, 0x55, 0x70, 0x73, 0x65, 0x72, 0x74, 0x55, 0x73, 0x65, 0x72, 0x50, 0x72, 0x65, 0x66, + 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x12, 0x39, 0x2e, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, + 0x72, 0x74, 0x2e, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, + 0x65, 0x73, 0x2e, 0x76, 0x31, 0x2e, 0x55, 0x70, 0x73, 0x65, 0x72, 0x74, 0x55, 0x73, 0x65, 0x72, + 0x50, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, + 0x73, 0x74, 0x1a, 0x16, 0x2e, 0x67, 0x6f, 0x6f, 0x67, 0x6c, 0x65, 0x2e, 0x70, 0x72, 0x6f, 0x74, + 0x6f, 0x62, 0x75, 0x66, 0x2e, 0x45, 0x6d, 0x70, 0x74, 0x79, 0x42, 0x59, 0x5a, 0x57, 0x67, 0x69, + 0x74, 0x68, 0x75, 0x62, 0x2e, 0x63, 0x6f, 0x6d, 0x2f, 0x67, 0x72, 0x61, 0x76, 0x69, 0x74, 0x61, + 0x74, 0x69, 0x6f, 0x6e, 0x61, 0x6c, 0x2f, 0x74, 0x65, 0x6c, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x2f, + 0x61, 0x70, 0x69, 0x2f, 0x67, 0x65, 0x6e, 0x2f, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x2f, 0x67, 0x6f, + 0x2f, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, 0x63, 0x65, 0x73, + 0x2f, 0x76, 0x31, 0x3b, 0x75, 0x73, 0x65, 0x72, 0x70, 0x72, 0x65, 0x66, 0x65, 0x72, 0x65, 0x6e, + 0x63, 0x65, 0x73, 0x76, 0x31, 0x62, 0x06, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x33, } var ( @@ -390,32 +374,30 @@ var file_teleport_userpreferences_v1_userpreferences_proto_goTypes = []interface (*GetUserPreferencesRequest)(nil), // 1: teleport.userpreferences.v1.GetUserPreferencesRequest (*GetUserPreferencesResponse)(nil), // 2: teleport.userpreferences.v1.GetUserPreferencesResponse (*UpsertUserPreferencesRequest)(nil), // 3: teleport.userpreferences.v1.UpsertUserPreferencesRequest - (*AssistUserPreferences)(nil), // 4: teleport.userpreferences.v1.AssistUserPreferences - (Theme)(0), // 5: teleport.userpreferences.v1.Theme - (*OnboardUserPreferences)(nil), // 6: teleport.userpreferences.v1.OnboardUserPreferences - (*ClusterUserPreferences)(nil), // 7: teleport.userpreferences.v1.ClusterUserPreferences - (*UnifiedResourcePreferences)(nil), // 8: teleport.userpreferences.v1.UnifiedResourcePreferences - (*AccessGraphUserPreferences)(nil), // 9: teleport.userpreferences.v1.AccessGraphUserPreferences - (*emptypb.Empty)(nil), // 10: google.protobuf.Empty + (Theme)(0), // 4: teleport.userpreferences.v1.Theme + (*OnboardUserPreferences)(nil), // 5: teleport.userpreferences.v1.OnboardUserPreferences + (*ClusterUserPreferences)(nil), // 6: teleport.userpreferences.v1.ClusterUserPreferences + (*UnifiedResourcePreferences)(nil), // 7: teleport.userpreferences.v1.UnifiedResourcePreferences + (*AccessGraphUserPreferences)(nil), // 8: teleport.userpreferences.v1.AccessGraphUserPreferences + (*emptypb.Empty)(nil), // 9: google.protobuf.Empty } var file_teleport_userpreferences_v1_userpreferences_proto_depIdxs = []int32{ - 4, // 0: teleport.userpreferences.v1.UserPreferences.assist:type_name -> teleport.userpreferences.v1.AssistUserPreferences - 5, // 1: teleport.userpreferences.v1.UserPreferences.theme:type_name -> teleport.userpreferences.v1.Theme - 6, // 2: teleport.userpreferences.v1.UserPreferences.onboard:type_name -> teleport.userpreferences.v1.OnboardUserPreferences - 7, // 3: teleport.userpreferences.v1.UserPreferences.cluster_preferences:type_name -> teleport.userpreferences.v1.ClusterUserPreferences - 8, // 4: teleport.userpreferences.v1.UserPreferences.unified_resource_preferences:type_name -> teleport.userpreferences.v1.UnifiedResourcePreferences - 9, // 5: teleport.userpreferences.v1.UserPreferences.access_graph:type_name -> teleport.userpreferences.v1.AccessGraphUserPreferences - 0, // 6: teleport.userpreferences.v1.GetUserPreferencesResponse.preferences:type_name -> teleport.userpreferences.v1.UserPreferences - 0, // 7: teleport.userpreferences.v1.UpsertUserPreferencesRequest.preferences:type_name -> teleport.userpreferences.v1.UserPreferences - 1, // 8: teleport.userpreferences.v1.UserPreferencesService.GetUserPreferences:input_type -> teleport.userpreferences.v1.GetUserPreferencesRequest - 3, // 9: teleport.userpreferences.v1.UserPreferencesService.UpsertUserPreferences:input_type -> teleport.userpreferences.v1.UpsertUserPreferencesRequest - 2, // 10: teleport.userpreferences.v1.UserPreferencesService.GetUserPreferences:output_type -> teleport.userpreferences.v1.GetUserPreferencesResponse - 10, // 11: teleport.userpreferences.v1.UserPreferencesService.UpsertUserPreferences:output_type -> google.protobuf.Empty - 10, // [10:12] is the sub-list for method output_type - 8, // [8:10] is the sub-list for method input_type - 8, // [8:8] is the sub-list for extension type_name - 8, // [8:8] is the sub-list for extension extendee - 0, // [0:8] is the sub-list for field type_name + 4, // 0: teleport.userpreferences.v1.UserPreferences.theme:type_name -> teleport.userpreferences.v1.Theme + 5, // 1: teleport.userpreferences.v1.UserPreferences.onboard:type_name -> teleport.userpreferences.v1.OnboardUserPreferences + 6, // 2: teleport.userpreferences.v1.UserPreferences.cluster_preferences:type_name -> teleport.userpreferences.v1.ClusterUserPreferences + 7, // 3: teleport.userpreferences.v1.UserPreferences.unified_resource_preferences:type_name -> teleport.userpreferences.v1.UnifiedResourcePreferences + 8, // 4: teleport.userpreferences.v1.UserPreferences.access_graph:type_name -> teleport.userpreferences.v1.AccessGraphUserPreferences + 0, // 5: teleport.userpreferences.v1.GetUserPreferencesResponse.preferences:type_name -> teleport.userpreferences.v1.UserPreferences + 0, // 6: teleport.userpreferences.v1.UpsertUserPreferencesRequest.preferences:type_name -> teleport.userpreferences.v1.UserPreferences + 1, // 7: teleport.userpreferences.v1.UserPreferencesService.GetUserPreferences:input_type -> teleport.userpreferences.v1.GetUserPreferencesRequest + 3, // 8: teleport.userpreferences.v1.UserPreferencesService.UpsertUserPreferences:input_type -> teleport.userpreferences.v1.UpsertUserPreferencesRequest + 2, // 9: teleport.userpreferences.v1.UserPreferencesService.GetUserPreferences:output_type -> teleport.userpreferences.v1.GetUserPreferencesResponse + 9, // 10: teleport.userpreferences.v1.UserPreferencesService.UpsertUserPreferences:output_type -> google.protobuf.Empty + 9, // [9:11] is the sub-list for method output_type + 7, // [7:9] is the sub-list for method input_type + 7, // [7:7] is the sub-list for extension type_name + 7, // [7:7] is the sub-list for extension extendee + 0, // [0:7] is the sub-list for field type_name } func init() { file_teleport_userpreferences_v1_userpreferences_proto_init() } @@ -424,7 +406,6 @@ func file_teleport_userpreferences_v1_userpreferences_proto_init() { return } file_teleport_userpreferences_v1_access_graph_proto_init() - file_teleport_userpreferences_v1_assist_proto_init() file_teleport_userpreferences_v1_cluster_preferences_proto_init() file_teleport_userpreferences_v1_onboard_proto_init() file_teleport_userpreferences_v1_theme_proto_init() diff --git a/api/proto/teleport/assist/v1/assist.proto b/api/proto/teleport/assist/v1/assist.proto deleted file mode 100644 index c5b6631a030..00000000000 --- a/api/proto/teleport/assist/v1/assist.proto +++ /dev/null @@ -1,201 +0,0 @@ -// Copyright 2023 Gravitational, Inc -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -syntax = "proto3"; - -package teleport.assist.v1; - -import "google/protobuf/empty.proto"; -import "google/protobuf/timestamp.proto"; -import "teleport/legacy/client/proto/authservice.proto"; - -option go_package = "github.com/gravitational/teleport/api/gen/proto/go/assist/v1;assist"; - -// GetAssistantMessagesRequest is a request to the assistant service. -message GetAssistantMessagesRequest { - // ConversationId identifies a conversation. - // It's used to tie all messages in a one conversation. - string conversation_id = 1; - // username is a username of the user who sent the message. - string username = 2; -} - -// AssistantMessage is a message sent to the assistant service. The conversation -// must be created first. -message AssistantMessage { - // type is a type of message. It can be Chat response/query or a command to run. - string type = 1; - // CreatedTime is the time when the event occurred. - google.protobuf.Timestamp created_time = 2; - // payload is a JSON message - string payload = 3; -} - -// CreateAssistantMessageRequest is a request to the assistant service. -message CreateAssistantMessageRequest { - // message is a message sent to the assistant service. - AssistantMessage message = 1; - // ConversationId is used to tie all messages into a conversation. - string conversation_id = 2; - // username is a username of the user who sent the message. - string username = 3; -} - -// GetAssistantMessagesResponse is a response from the assistant service. -message GetAssistantMessagesResponse { - // messages is a list of messages. - repeated AssistantMessage messages = 1; -} - -// GetAssistantConversationsRequest is a request to get a list of conversations. -message GetAssistantConversationsRequest { - // username is a username of the user who created the conversation. - string username = 1; -} - -// ConversationInfo is a conversation info. It contains a conversation -// information like ID, title, created time. -message ConversationInfo { - // id is a unique conversation ID. - string id = 1; - // title is a title of the conversation. - string title = 2; - // createdTime is the time when the conversation was created. - google.protobuf.Timestamp created_time = 3; -} - -// GetAssistantConversationsResponse is a response from the assistant service. -message GetAssistantConversationsResponse { - // conversations is a list of conversations. - repeated ConversationInfo conversations = 1; -} - -// CreateAssistantConversationRequest is a request to create a new conversation. -message CreateAssistantConversationRequest { - // username is a username of the user who created the conversation. - string username = 1; - // createdTime is the time when the conversation was created. - google.protobuf.Timestamp created_time = 2; -} - -// CreateAssistantConversationResponse is a response from the assistant service. -message CreateAssistantConversationResponse { - // id is a unique conversation ID. - string id = 1; -} - -// UpdateAssistantConversationInfoRequest is a request to update the conversation info. -message UpdateAssistantConversationInfoRequest { - // conversationId is a unique conversation ID. - string conversation_id = 1; - // username is a username of the user who created the conversation. - string username = 2; - // title is a title of the conversation. - string title = 3; -} - -// IsAssistEnabledRequest is a request to the assistant service on if assist is enabled or not. -message IsAssistEnabledRequest {} - -// IsAssistEnabledResponse is a response from the assistant service on if assist is enabled or not. -message IsAssistEnabledResponse { - // enabled is true if the assist is enabled or not on the auth level. - bool enabled = 1; -} - -// DeleteAssistantConversationRequest is a request to delete the conversation. -message DeleteAssistantConversationRequest { - // conversationId is a unique conversation ID. - string conversation_id = 1; - // username is a username of the user who created the conversation. - string username = 2; -} - -// GetAssistantEmbeddingsRequest is a request to get embeddings. -message GetAssistantEmbeddingsRequest { - // username is a username of the user who requested the embeddings. - string username = 1; - // query is the query used for similarity search. - string query = 2; - // limit is the number of embeddings to return (also known as k). - uint32 limit = 3; - // kind is the kind of embeddings to return (ex, node). - string kind = 4; -} - -// EmbeddingDocument is a document with an embedding. -message EmbeddedDocument { - // id is the id of the document. - string id = 1; - // content is the content of the document. - string content = 2; - // similarityScore is the similarity score of the document. - float similarity_score = 3; -} - -// GetAssistantEmbeddingsResponse is a response from the assistant service. -message GetAssistantEmbeddingsResponse { - // embeddings is the list of embeddings. - // The list is sorted by similarity score in descending order. - repeated EmbeddedDocument embeddings = 1; -} - -// SearchUnifiedResourcesRequest is a request to search for one or more resource kinds using similiarity search. -message SearchUnifiedResourcesRequest { - // query is the query used for similarity search. - string query = 1; - // limit is the number of embeddings to return (also known as k). - int32 limit = 2; - // kinds is the kind of embeddings to return (ex, node). Returns all supported kinds if empty. - repeated string kinds = 3; -} - -// SearchUnifiedResourcesResponse is a response from the assistant service with a similarity-ordered list of resources. -message SearchUnifiedResourcesResponse { - // resources is the list of resources. - repeated proto.PaginatedResource resources = 1; -} - -// AssistService is a service that provides an ability to communicate with the Teleport Assist. -service AssistService { - // CreateNewConversation creates a new conversation and returns the UUID of it. - rpc CreateAssistantConversation(CreateAssistantConversationRequest) returns (CreateAssistantConversationResponse); - - // GetAssistantConversations returns all conversations for the connected user. - rpc GetAssistantConversations(GetAssistantConversationsRequest) returns (GetAssistantConversationsResponse); - - // DeleteAssistantConversation deletes the conversation and all messages associated with it. - rpc DeleteAssistantConversation(DeleteAssistantConversationRequest) returns (google.protobuf.Empty); - - // GetAssistantMessages returns all messages associated with the given conversation ID. - rpc GetAssistantMessages(GetAssistantMessagesRequest) returns (GetAssistantMessagesResponse); - - // CreateAssistantMessage creates a new message in the given conversation. - rpc CreateAssistantMessage(CreateAssistantMessageRequest) returns (google.protobuf.Empty); - - // UpdateAssistantConversationInfo updates the conversation info. - rpc UpdateAssistantConversationInfo(UpdateAssistantConversationInfoRequest) returns (google.protobuf.Empty); - - // IsAssistEnabled returns true if the assist is enabled or not on the auth level. - rpc IsAssistEnabled(IsAssistEnabledRequest) returns (IsAssistEnabledResponse); - - // SearchUnifiedResources returns a similarity-ordered list of resources from the unified resource cache. - rpc SearchUnifiedResources(SearchUnifiedResourcesRequest) returns (SearchUnifiedResourcesResponse); -} - -// AssistEmbeddingService is a service that provides an ability to communicate with the Assist Embedding service. -service AssistEmbeddingService { - // AssistantGetEmbeddings returns the embeddings for the given query. - rpc GetAssistantEmbeddings(GetAssistantEmbeddingsRequest) returns (GetAssistantEmbeddingsResponse); -} diff --git a/api/proto/teleport/userpreferences/v1/userpreferences.proto b/api/proto/teleport/userpreferences/v1/userpreferences.proto index 7ef8e0f980f..7537cbc21ea 100644 --- a/api/proto/teleport/userpreferences/v1/userpreferences.proto +++ b/api/proto/teleport/userpreferences/v1/userpreferences.proto @@ -18,7 +18,6 @@ package teleport.userpreferences.v1; import "google/protobuf/empty.proto"; import "teleport/userpreferences/v1/access_graph.proto"; -import "teleport/userpreferences/v1/assist.proto"; import "teleport/userpreferences/v1/cluster_preferences.proto"; import "teleport/userpreferences/v1/onboard.proto"; import "teleport/userpreferences/v1/theme.proto"; @@ -29,7 +28,8 @@ option go_package = "github.com/gravitational/teleport/api/gen/proto/go/userpref // UserPreferences is a collection of different user changeable preferences for the frontend. message UserPreferences { // assist is the preferences for the Teleport Assist. - v1.AssistUserPreferences assist = 1; + reserved 1; + reserved "assist"; // theme is the theme of the frontend. Theme theme = 2; // onboard is the preferences from the onboarding questionnaire. diff --git a/api/types/constants.go b/api/types/constants.go index a217e3f6801..c41c7d8c766 100644 --- a/api/types/constants.go +++ b/api/types/constants.go @@ -463,10 +463,6 @@ const ( // KindHeadlessAuthentication is a headless authentication resource. KindHeadlessAuthentication = "headless_authentication" - // KindAssistant is used to program RBAC for - // Teleport Assist resources. - KindAssistant = "assistant" - // KindAccessGraph is the RBAC kind for access graph. KindAccessGraph = "access_graph" diff --git a/api/types/networking.go b/api/types/networking.go index d0f82ea48a6..831ef7a6503 100644 --- a/api/types/networking.go +++ b/api/types/networking.go @@ -107,12 +107,6 @@ type ClusterNetworkingConfig interface { // SetProxyPingInterval sets the proxy ping interval. SetProxyPingInterval(time.Duration) - // GetAssistCommandExecutionWorkers gets the number of parallel command execution workers for Assist - GetAssistCommandExecutionWorkers() int32 - - // SetAssistCommandExecutionWorkers sets the number of parallel command execution workers for Assist - SetAssistCommandExecutionWorkers(n int32) - // GetCaseInsensitiveRouting gets the case-insensitive routing option. GetCaseInsensitiveRouting() bool @@ -368,12 +362,6 @@ func (c *ClusterNetworkingConfigV2) CheckAndSetDefaults() error { return trace.Wrap(err) } - if c.Spec.AssistCommandExecutionWorkers < 0 { - return trace.BadParameter("command_execution_workers must be non-negative") - } else if c.Spec.AssistCommandExecutionWorkers == 0 { - c.Spec.AssistCommandExecutionWorkers = defaults.AssistCommandExecutionWorkers - } - return nil } @@ -387,16 +375,6 @@ func (c *ClusterNetworkingConfigV2) SetProxyPingInterval(interval time.Duration) c.Spec.ProxyPingInterval = Duration(interval) } -// GetAssistCommandExecutionWorkers gets the number of parallel command execution workers for Assist -func (c *ClusterNetworkingConfigV2) GetAssistCommandExecutionWorkers() int32 { - return c.Spec.AssistCommandExecutionWorkers -} - -// SetAssistCommandExecutionWorkers sets the number of parallel command execution workers for Assist -func (c *ClusterNetworkingConfigV2) SetAssistCommandExecutionWorkers(n int32) { - c.Spec.AssistCommandExecutionWorkers = n -} - // GetCaseInsensitiveRouting gets the case-insensitive routing option. func (c *ClusterNetworkingConfigV2) GetCaseInsensitiveRouting() bool { return c.Spec.CaseInsensitiveRouting diff --git a/constants.go b/constants.go index 48edeeeb55a..fe3b1914705 100644 --- a/constants.go +++ b/constants.go @@ -277,13 +277,10 @@ const ( // ComponentAthena represents athena clients. ComponentAthena = "athena" - // ComponentProxySecureGRPC represents secure gRPC server running on Proxy (used for Kube). + // ComponentProxySecureGRPC represents a secure gRPC server running on Proxy (used for Kube). ComponentProxySecureGRPC = "proxy:secure-grpc" - // ComponentAssist represents Teleport Assist - ComponentAssist = "assist" - - // VerboseLogEnvVar forces all logs to be verbose (down to DEBUG level) + // VerboseLogsEnvVar forces all logs to be verbose (down to DEBUG level) VerboseLogsEnvVar = "TELEPORT_DEBUG" // IterationsEnvVar sets tests iterations to run diff --git a/gen/proto/ts/teleport/userpreferences/v1/userpreferences_pb.ts b/gen/proto/ts/teleport/userpreferences/v1/userpreferences_pb.ts index 04043998bbd..dddf6e429f8 100644 --- a/gen/proto/ts/teleport/userpreferences/v1/userpreferences_pb.ts +++ b/gen/proto/ts/teleport/userpreferences/v1/userpreferences_pb.ts @@ -34,19 +34,12 @@ import { UnifiedResourcePreferences } from "./unified_resource_preferences_pb"; import { ClusterUserPreferences } from "./cluster_preferences_pb"; import { OnboardUserPreferences } from "./onboard_pb"; import { Theme } from "./theme_pb"; -import { AssistUserPreferences } from "./assist_pb"; /** * UserPreferences is a collection of different user changeable preferences for the frontend. * * @generated from protobuf message teleport.userpreferences.v1.UserPreferences */ export interface UserPreferences { - /** - * assist is the preferences for the Teleport Assist. - * - * @generated from protobuf field: teleport.userpreferences.v1.AssistUserPreferences assist = 1; - */ - assist?: AssistUserPreferences; /** * theme is the theme of the frontend. * @@ -115,7 +108,6 @@ export interface UpsertUserPreferencesRequest { class UserPreferences$Type extends MessageType { constructor() { super("teleport.userpreferences.v1.UserPreferences", [ - { no: 1, name: "assist", kind: "message", T: () => AssistUserPreferences }, { no: 2, name: "theme", kind: "enum", T: () => ["teleport.userpreferences.v1.Theme", Theme, "THEME_"] }, { no: 3, name: "onboard", kind: "message", T: () => OnboardUserPreferences }, { no: 4, name: "cluster_preferences", kind: "message", T: () => ClusterUserPreferences }, @@ -135,9 +127,6 @@ class UserPreferences$Type extends MessageType { while (reader.pos < end) { let [fieldNo, wireType] = reader.tag(); switch (fieldNo) { - case /* teleport.userpreferences.v1.AssistUserPreferences assist */ 1: - message.assist = AssistUserPreferences.internalBinaryRead(reader, reader.uint32(), options, message.assist); - break; case /* teleport.userpreferences.v1.Theme theme */ 2: message.theme = reader.int32(); break; @@ -165,9 +154,6 @@ class UserPreferences$Type extends MessageType { return message; } internalBinaryWrite(message: UserPreferences, writer: IBinaryWriter, options: BinaryWriteOptions): IBinaryWriter { - /* teleport.userpreferences.v1.AssistUserPreferences assist = 1; */ - if (message.assist) - AssistUserPreferences.internalBinaryWrite(message.assist, writer.tag(1, WireType.LengthDelimited).fork(), options).join(); /* teleport.userpreferences.v1.Theme theme = 2; */ if (message.theme !== 0) writer.tag(2, WireType.Varint).int32(message.theme); diff --git a/go.mod b/go.mod index 7ae09d7c16c..db5a402d99a 100644 --- a/go.mod +++ b/go.mod @@ -167,7 +167,6 @@ require ( github.com/redis/go-redis/v9 v9.5.1 // replaced github.com/russellhaering/gosaml2 v0.9.1 github.com/russellhaering/goxmldsig v1.4.0 - github.com/sashabaranov/go-openai v1.23.0 github.com/schollz/progressbar/v3 v3.14.2 github.com/scim2/filter-parser/v2 v2.2.0 github.com/segmentio/parquet-go v0.0.0-20230712180008-5d42db8f0d47 diff --git a/go.sum b/go.sum index 16aa0b68586..ed98fd77db3 100644 --- a/go.sum +++ b/go.sum @@ -2130,8 +2130,6 @@ github.com/sagikazarmark/slog-shim v0.1.0 h1:diDBnUNK9N/354PgrxMywXnAwEr1QZcOr6g github.com/sagikazarmark/slog-shim v0.1.0/go.mod h1:SrcSrq8aKtyuqEI1uvTDTK1arOWRIczQRv+GVI1AkeQ= github.com/sasha-s/go-deadlock v0.3.1 h1:sqv7fDNShgjcaxkO0JNcOAlr8B9+cV5Ey/OB71efZx0= github.com/sasha-s/go-deadlock v0.3.1/go.mod h1:F73l+cr82YSh10GxyRI6qZiCgK64VaZjwesgfQ1/iLM= -github.com/sashabaranov/go-openai v1.23.0 h1:KYW97r5yc35PI2MxeLZ3OofecB/6H+yxvSNqiT9u8is= -github.com/sashabaranov/go-openai v1.23.0/go.mod h1:lj5b/K+zjTSFxVLijLSTDZuP7adOgerWeFyZLUhAKRg= github.com/sassoftware/relic v7.2.1+incompatible h1:Pwyh1F3I0r4clFJXkSI8bOyJINGqpgjJU3DYAZeI05A= github.com/sassoftware/relic v7.2.1+incompatible/go.mod h1:CWfAxv73/iLZ17rbyhIEq3K9hs5w6FpNMdUT//qR+zk= github.com/sassoftware/relic/v7 v7.6.2 h1:rS44Lbv9G9eXsukknS4mSjIAuuX+lMq/FnStgmZlUv4= diff --git a/integrations/terraform/tfschema/types_terraform.go b/integrations/terraform/tfschema/types_terraform.go index 3b052a5953a..553694f1d5b 100644 --- a/integrations/terraform/tfschema/types_terraform.go +++ b/integrations/terraform/tfschema/types_terraform.go @@ -958,11 +958,6 @@ func GenSchemaClusterNetworkingConfigV2(ctx context.Context) (github_com_hashico }, "spec": { Attributes: github_com_hashicorp_terraform_plugin_framework_tfsdk.SingleNestedAttributes(map[string]github_com_hashicorp_terraform_plugin_framework_tfsdk.Attribute{ - "assist_command_execution_workers": { - Description: "AssistCommandExecutionWorkers determines the number of workers that will execute arbitrary Assist commands on servers in parallel", - Optional: true, - Type: github_com_hashicorp_terraform_plugin_framework_types.Int64Type, - }, "case_insensitive_routing": { Description: "CaseInsensitiveRouting causes proxies to use case-insensitive hostname matching.", Optional: true, @@ -10817,23 +10812,6 @@ func CopyClusterNetworkingConfigV2FromTerraform(_ context.Context, tf github_com } } } - { - a, ok := tf.Attrs["assist_command_execution_workers"] - if !ok { - diags.Append(attrReadMissingDiag{"ClusterNetworkingConfigV2.Spec.AssistCommandExecutionWorkers"}) - } else { - v, ok := a.(github_com_hashicorp_terraform_plugin_framework_types.Int64) - if !ok { - diags.Append(attrReadConversionFailureDiag{"ClusterNetworkingConfigV2.Spec.AssistCommandExecutionWorkers", "github.com/hashicorp/terraform-plugin-framework/types.Int64"}) - } else { - var t int32 - if !v.Null && !v.Unknown { - t = int32(v.Value) - } - obj.AssistCommandExecutionWorkers = t - } - } - } { a, ok := tf.Attrs["case_insensitive_routing"] if !ok { @@ -11473,28 +11451,6 @@ func CopyClusterNetworkingConfigV2ToTerraform(ctx context.Context, obj *github_c tf.Attrs["proxy_ping_interval"] = v } } - { - t, ok := tf.AttrTypes["assist_command_execution_workers"] - if !ok { - diags.Append(attrWriteMissingDiag{"ClusterNetworkingConfigV2.Spec.AssistCommandExecutionWorkers"}) - } else { - v, ok := tf.Attrs["assist_command_execution_workers"].(github_com_hashicorp_terraform_plugin_framework_types.Int64) - if !ok { - i, err := t.ValueFromTerraform(ctx, github_com_hashicorp_terraform_plugin_go_tftypes.NewValue(t.TerraformType(ctx), nil)) - if err != nil { - diags.Append(attrWriteGeneralError{"ClusterNetworkingConfigV2.Spec.AssistCommandExecutionWorkers", err}) - } - v, ok = i.(github_com_hashicorp_terraform_plugin_framework_types.Int64) - if !ok { - diags.Append(attrWriteConversionFailureDiag{"ClusterNetworkingConfigV2.Spec.AssistCommandExecutionWorkers", "github.com/hashicorp/terraform-plugin-framework/types.Int64"}) - } - v.Null = int64(obj.AssistCommandExecutionWorkers) == 0 - } - v.Value = int64(obj.AssistCommandExecutionWorkers) - v.Unknown = false - tf.Attrs["assist_command_execution_workers"] = v - } - } { t, ok := tf.AttrTypes["case_insensitive_routing"] if !ok { diff --git a/lib/ai/chat.go b/lib/ai/chat.go deleted file mode 100644 index 73ed68d982e..00000000000 --- a/lib/ai/chat.go +++ /dev/null @@ -1,92 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package ai - -import ( - "context" - - "github.com/gravitational/trace" - "github.com/sashabaranov/go-openai" - - "github.com/gravitational/teleport/lib/ai/model" - "github.com/gravitational/teleport/lib/ai/model/output" - "github.com/gravitational/teleport/lib/ai/tokens" -) - -// Chat represents a conversation between a user and an assistant with context memory. -type Chat struct { - client *Client - messages []openai.ChatCompletionMessage - agent *model.Agent -} - -// Insert inserts a message into the conversation. Returns the index of the message. -func (chat *Chat) Insert(role string, content string) int { - chat.messages = append(chat.messages, openai.ChatCompletionMessage{ - Role: role, - Content: content, - }) - - return len(chat.messages) - 1 -} - -// GetMessages returns the messages in the conversation. -func (chat *Chat) GetMessages() []openai.ChatCompletionMessage { - return chat.messages -} - -// Complete completes the conversation with a message from the assistant based on the current context and user input. -// On success, it returns the message. -// Returned types: -// - message: one of the message types below -// - error: an error if one occurred -// Message types: -// - CompletionCommand: a command from the assistant -// - Message: a text message from the assistant -// - AccessRequest: an access request suggestion from the assistant -func (chat *Chat) Complete(ctx context.Context, userInput string, progressUpdates func(*model.AgentAction)) (any, *tokens.TokenCount, error) { - // if the chat is empty, return the initial response we predefine instead of querying GPT-4 - if len(chat.messages) == 1 { - return &output.Message{ - Content: model.InitialAIResponse, - }, tokens.NewTokenCount(), nil - } - - return chat.Reply(ctx, userInput, progressUpdates) -} - -// Reply replies to the user input with a message from the assistant based on the current context. -func (chat *Chat) Reply(ctx context.Context, userInput string, progressUpdates func(*model.AgentAction)) (any, *tokens.TokenCount, error) { - userMessage := openai.ChatCompletionMessage{ - Role: openai.ChatMessageRoleUser, - Content: userInput, - } - - response, tokenCount, err := chat.agent.PlanAndExecute(ctx, chat.client.svc, chat.messages, userMessage, progressUpdates) - if err != nil { - return nil, nil, trace.Wrap(err) - } - - return response, tokenCount, nil -} - -// Clear clears the conversation. -func (chat *Chat) Clear() { - chat.messages = []openai.ChatCompletionMessage{} -} diff --git a/lib/ai/chat_test.go b/lib/ai/chat_test.go deleted file mode 100644 index bf0b1bf5394..00000000000 --- a/lib/ai/chat_test.go +++ /dev/null @@ -1,364 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package ai - -import ( - "context" - "encoding/json" - "fmt" - "net/http" - "net/http/httptest" - "testing" - - "github.com/sashabaranov/go-openai" - "github.com/stretchr/testify/require" - - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/lib/ai/model" - "github.com/gravitational/teleport/lib/ai/model/output" - "github.com/gravitational/teleport/lib/ai/model/tools" - "github.com/gravitational/teleport/lib/ai/testutils" - "github.com/gravitational/teleport/lib/modules" -) - -func TestChat_PromptTokens(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - messages []openai.ChatCompletionMessage - want int - }{ - { - name: "empty", - messages: []openai.ChatCompletionMessage{}, - want: 0, - }, - { - name: "only system message", - messages: []openai.ChatCompletionMessage{ - { - Role: openai.ChatMessageRoleSystem, - Content: "Hello", - }, - }, - want: 850, - }, - { - name: "system and user messages", - messages: []openai.ChatCompletionMessage{ - { - Role: openai.ChatMessageRoleSystem, - Content: "Hello", - }, - { - Role: openai.ChatMessageRoleUser, - Content: "Hi LLM.", - }, - }, - want: 855, - }, - { - name: "tokenize our prompt", - messages: []openai.ChatCompletionMessage{ - { - Role: openai.ChatMessageRoleSystem, - Content: model.PromptCharacter("Bob"), - }, - { - Role: openai.ChatMessageRoleUser, - Content: "Show me free disk space on localhost node.", - }, - }, - want: 1114, - }, - } - - for _, tt := range tests { - tt := tt - t.Run(tt.name, func(t *testing.T) { - responses := []string{ - generateCommandResponse(t), - } - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - - require.NotEmpty(t, responses, "Unexpected request") - dataBytes := responses[0] - _, err := w.Write([]byte(dataBytes)) - require.NoError(t, err, "Write error") - - responses = responses[1:] - })) - - t.Cleanup(server.Close) - - cfg := openai.DefaultConfig("secret-test-token") - cfg.BaseURL = server.URL + "/v1" - - client := NewClientFromConfig(cfg) - - toolContext := tools.ToolContext{ - User: "Bob", - } - chat := client.NewChat(&toolContext) - - for _, message := range tt.messages { - chat.Insert(message.Role, message.Content) - } - - ctx := context.Background() - _, tokenCount, err := chat.Complete(ctx, "", func(aa *model.AgentAction) {}) - require.NoError(t, err) - - prompt, completion := tokenCount.CountAll() - usedTokens := prompt + completion - require.Equal(t, tt.want, usedTokens) - }) - } -} - -func TestChat_Complete(t *testing.T) { - beforeModules := modules.GetModules() - modules.SetModules(&modules.TestModules{ - TestBuildType: modules.BuildEnterprise, - }) - t.Cleanup(func() { modules.SetModules(beforeModules) }) - - responses := [][]byte{ - []byte(generateTextResponse()), - []byte(generateCommandResponse(t)), - []byte(generateAccessRequestResponse(t)), - } - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - - require.NotEmpty(t, responses, "Unexpected request") - dataBytes := responses[0] - - _, err := w.Write(dataBytes) - require.NoError(t, err, "Write error") - - responses = responses[1:] - })) - defer server.Close() - - cfg := openai.DefaultConfig("secret-test-token") - cfg.BaseURL = server.URL + "/v1" - client := NewClientFromConfig(cfg) - - toolContext := tools.ToolContext{ - User: "Bob", - } - chat := client.NewChat(&toolContext) - - ctx := context.Background() - _, _, err := chat.Complete(ctx, "Hello", func(aa *model.AgentAction) {}) - require.NoError(t, err) - - chat.Insert(openai.ChatMessageRoleUser, "Show me free disk space on localhost node.") - - t.Run("text completion", func(t *testing.T) { - msg, _, err := chat.Complete(ctx, "Show me free disk space", func(aa *model.AgentAction) {}) - require.NoError(t, err) - - require.IsType(t, &output.StreamingMessage{}, msg) - streamingMessage := msg.(*output.StreamingMessage) - require.Equal(t, "Which ", <-streamingMessage.Parts) - require.Equal(t, "node do ", <-streamingMessage.Parts) - require.Equal(t, "you want ", <-streamingMessage.Parts) - require.Equal(t, "use?", <-streamingMessage.Parts) - }) - - t.Run("command completion", func(t *testing.T) { - msg, _, err := chat.Complete(ctx, "localhost", func(aa *model.AgentAction) {}) - require.NoError(t, err) - - require.IsType(t, &output.CompletionCommand{}, msg) - command := msg.(*output.CompletionCommand) - require.Equal(t, "df -h", command.Command) - require.Len(t, command.Nodes, 1) - require.Equal(t, "localhost", command.Nodes[0]) - }) - - t.Run("access request creation", func(t *testing.T) { - msg, _, err := chat.Complete(ctx, "Now, request access to the resource with kind node, hostname Alpha.local and the Name a35161f0-a2dc-48e7-bdd2-49b81926cab7", func(aa *model.AgentAction) {}) - require.NoError(t, err) - - require.IsType(t, &output.AccessRequest{}, msg) - request := msg.(*output.AccessRequest) - require.Empty(t, request.Roles) - require.Empty(t, request.SuggestedReviewers) - require.Equal(t, "maintenance", request.Reason) - require.Len(t, request.Resources, 1) - require.Equal(t, "a35161f0-a2dc-48e7-bdd2-49b81926cab7", request.Resources[0].Name) - require.Equal(t, "Alpha.local", request.Resources[0].FriendlyName) - }) -} - -// generateTextResponse generates a response for a text completion -func generateTextResponse() string { - dataBytes := []byte{} - dataBytes = append(dataBytes, []byte("event: message\n")...) - - data := `{"id":"1","object":"completion","created":1598069254,"model":"gpt-4","choices":[{"index": 0, "delta":{"content": "Which ", "role": "assistant"}}]}` - dataBytes = append(dataBytes, []byte("data: "+data+"\n\n")...) - dataBytes = append(dataBytes, []byte("event: message\n")...) - - data = `{"id":"2","object":"completion","created":1598069254,"model":"gpt-4","choices":[{"index": 0, "delta":{"content": "node do ", "role": "assistant"}}]}` - dataBytes = append(dataBytes, []byte("data: "+data+"\n\n")...) - dataBytes = append(dataBytes, []byte("event: message\n")...) - - data = `{"id":"3","object":"completion","created":1598069255,"model":"gpt-4","choices":[{"index": 0, "delta":{"content": "you want ", "role": "assistant"}}]}` - dataBytes = append(dataBytes, []byte("data: "+data+"\n\n")...) - dataBytes = append(dataBytes, []byte("event: message\n")...) - - data = `{"id":"4","object":"completion","created":1598069254,"model":"gpt-4","choices":[{"index": 0, "delta":{"content": "use?", "role": "assistant"}}]}` - dataBytes = append(dataBytes, []byte("data: "+data+"\n\n")...) - dataBytes = append(dataBytes, []byte("event: done\n")...) - - dataBytes = append(dataBytes, []byte("data: [DONE]\n\n")...) - - return string(dataBytes) -} - -// generateCommandResponse generates a response for the command "df -h" on the node "localhost" -func generateCommandResponse(t *testing.T) string { - dataBytes := []byte{} - dataBytes = append(dataBytes, []byte("event: message\n")...) - - actionObj := model.PlanOutput{ - Action: "Command Execution", - ActionInput: struct { - Command string `json:"command"` - Nodes []string `json:"nodes"` - }{"df -h", []string{"localhost"}}, - } - actionJson, err := json.Marshal(actionObj) - if err != nil { - require.NoError(t, err) - } - - obj := struct { - Content string `json:"content"` - Role string `json:"role"` - }{ - Content: string(actionJson), - Role: "assistant", - } - json, err := json.Marshal(obj) - if err != nil { - require.NoError(t, err) - } - - data := fmt.Sprintf(`{"id":"1","object":"completion","created":1598069254,"model":"gpt-4","choices":[{"index": 0, "delta":%v}]}`, string(json)) - dataBytes = append(dataBytes, []byte("data: "+data+"\n\n")...) - - dataBytes = append(dataBytes, []byte("event: done\n")...) - dataBytes = append(dataBytes, []byte("data: [DONE]\n\n")...) - - return string(dataBytes) -} - -func generateAccessRequestResponse(t *testing.T) string { - dataBytes := []byte{} - dataBytes = append(dataBytes, []byte("event: message\n")...) - - actionObj := model.PlanOutput{ - Action: "Create Access Requests", - ActionInput: struct { - SuggestedReviewers []string `json:"suggested_reviewers"` - Roles []string `json:"roles"` - Resources []output.Resource `json:"resources"` - Reason string `json:"reason"` - }{ - nil, - nil, - []output.Resource{{Type: types.KindNode, Name: "a35161f0-a2dc-48e7-bdd2-49b81926cab7", FriendlyName: "Alpha.local"}}, - "maintenance", - }, - } - actionJson, err := json.Marshal(actionObj) - if err != nil { - require.NoError(t, err) - } - - obj := struct { - Content string `json:"content"` - Role string `json:"role"` - }{ - Content: string(actionJson), - Role: "assistant", - } - json, err := json.Marshal(obj) - if err != nil { - require.NoError(t, err) - } - - data := fmt.Sprintf(`{"id":"1","object":"completion","created":1598069254,"model":"gpt-4","choices":[{"index": 0, "delta":%v}]}`, string(json)) - dataBytes = append(dataBytes, []byte("data: "+data+"\n\n")...) - - dataBytes = append(dataBytes, []byte("event: done\n")...) - dataBytes = append(dataBytes, []byte("data: [DONE]\n\n")...) - - return string(dataBytes) -} - -func TestChat_Complete_AuditQuery(t *testing.T) { - // Test setup: generate the responses that will be served by our OpenAI mock - action := model.PlanOutput{ - Action: tools.AuditQueryGenerationToolName, - ActionInput: "Lists user who connected to a server as root.", - Reasoning: "foo", - } - selectedAction, err := json.Marshal(action) - require.NoError(t, err) - const generatedQuery = "SELECT user FROM session_start WHERE login='root'" - - responses := []string{ - // The model must select the audit query tool - string(selectedAction), - // Then the audit query tool chooses to request session.start events - "session.start", - // Finally the tool builds a query based on the provided schemas - generatedQuery, - } - server := httptest.NewServer(testutils.GetTestHandlerFn(t, responses)) - t.Cleanup(server.Close) - - cfg := openai.DefaultConfig("secret-test-token") - cfg.BaseURL = server.URL - - client := NewClientFromConfig(cfg) - - // End of test setup, we run the agent - chat := client.NewAuditQuery("bob") - - ctx := context.Background() - // We insert a message to make the conversation not empty and skip the - // greeting message. - chat.Insert(openai.ChatMessageRoleUser, "Hello") - result, _, err := chat.Complete(ctx, "List users who connected to a server as root", func(action *model.AgentAction) {}) - require.NoError(t, err) - - // We check that the agent returns the expected response - message, ok := result.(*output.StreamingMessage) - require.True(t, ok) - require.Equal(t, generatedQuery, message.WaitAndConsume()) -} diff --git a/lib/ai/client.go b/lib/ai/client.go deleted file mode 100644 index 732c6be7c4f..00000000000 --- a/lib/ai/client.go +++ /dev/null @@ -1,253 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package ai - -import ( - "context" - - "github.com/gravitational/trace" - "github.com/sashabaranov/go-openai" - - "github.com/gravitational/teleport/lib/ai/embedding" - "github.com/gravitational/teleport/lib/ai/model" - modeltools "github.com/gravitational/teleport/lib/ai/model/tools" - "github.com/gravitational/teleport/lib/ai/tokens" - "github.com/gravitational/teleport/lib/modules" -) - -const ( - maxOpenAIEmbeddingsPerRequest = 1000 -) - -// Client is a client for OpenAI API. -type Client struct { - svc *openai.Client -} - -// NewClient creates a new client for OpenAI API. -func NewClient(authToken string) *Client { - return &Client{openai.NewClient(authToken)} -} - -// NewClientFromConfig creates a new client for OpenAI API from config. -func NewClientFromConfig(config openai.ClientConfig) *Client { - return &Client{openai.NewClientWithConfig(config)} -} - -// NewChat creates a new chat. The username is set in the conversation context, -// so that the AI can use it to personalize the conversation. -// embeddingServiceClient is used to get the embeddings from the Auth Server. -func (client *Client) NewChat(toolContext *modeltools.ToolContext) *Chat { - tools := []modeltools.Tool{ - &modeltools.CommandExecutionTool{}, - &modeltools.EmbeddingRetrievalTool{}, - } - - // The following tools are only available in the enterprise build. They will fail - // if included in OSS due to the lack of the required backend APIs. - if modules.GetModules().BuildType() == modules.BuildEnterprise { - tools = append(tools, &modeltools.AccessRequestCreateTool{}, - &modeltools.AccessRequestsListTool{}, - &modeltools.AccessRequestListRequestableRolesTool{}, - &modeltools.AccessRequestListRequestableResourcesTool{}) - } - - return &Chat{ - client: client, - messages: []openai.ChatCompletionMessage{ - { - Role: openai.ChatMessageRoleSystem, - Content: model.PromptCharacter(toolContext.User), - }, - }, - agent: model.NewAgent(toolContext, tools...), - } -} - -func (client *Client) NewCommand(username string) *Chat { - toolContext := &modeltools.ToolContext{User: username} - return &Chat{ - client: client, - messages: []openai.ChatCompletionMessage{ - { - Role: openai.ChatMessageRoleSystem, - Content: model.PromptCharacter(username), - }, - }, - agent: model.NewAgent(toolContext, &modeltools.CommandGenerationTool{}), - } -} - -func (client *Client) RunTool(ctx context.Context, toolContext *modeltools.ToolContext, toolName, toolInput string) (any, *tokens.TokenCount, error) { - tools := []modeltools.Tool{ - &modeltools.CommandExecutionTool{}, - &modeltools.EmbeddingRetrievalTool{}, - &modeltools.AuditQueryGenerationTool{LLM: client.svc}, - } - // The following tools are only available in the enterprise build. They will fail - // if included in OSS due to the lack of the required backend APIs. - if modules.GetModules().BuildType() == modules.BuildEnterprise { - tools = append(tools, &modeltools.AccessRequestCreateTool{}, - &modeltools.AccessRequestsListTool{}, - &modeltools.AccessRequestListRequestableRolesTool{}, - &modeltools.AccessRequestListRequestableResourcesTool{}) - } - agent := model.NewAgent(toolContext, tools...) - action := &model.AgentAction{ - Action: toolName, - Input: toolInput, - Reasoning: "Tool invoked directly", - } - - return agent.DoAction(ctx, client.svc, action) -} - -func (client *Client) NewAuditQuery(username string) *Chat { - toolContext := &modeltools.ToolContext{User: username} - return &Chat{ - client: client, - messages: []openai.ChatCompletionMessage{ - { - Role: openai.ChatMessageRoleSystem, - Content: model.PromptCharacter(username), - }, - }, - agent: model.NewAgent(toolContext, &modeltools.AuditQueryGenerationTool{LLM: client.svc}), - } -} - -// Summary creates a short summary for the given input. -func (client *Client) Summary(ctx context.Context, message string) (string, error) { - resp, err := client.svc.CreateChatCompletion( - ctx, - openai.ChatCompletionRequest{ - Model: openai.GPT4, - Messages: []openai.ChatCompletionMessage{ - {Role: openai.ChatMessageRoleSystem, Content: model.PromptSummarizeTitle}, - {Role: openai.ChatMessageRoleUser, Content: message}, - }, - }, - ) - - if err != nil { - return "", trace.Wrap(err) - } - - return resp.Choices[0].Message.Content, nil -} - -// CommandSummary creates a command summary based on the command output. -// The message history is also passed to the model in order to keep context -// and extract relevant information from the output. -func (client *Client) CommandSummary(ctx context.Context, messages []openai.ChatCompletionMessage, output map[string][]byte) (string, *tokens.TokenCount, error) { - messages = append(messages, openai.ChatCompletionMessage{ - Role: openai.ChatMessageRoleUser, Content: model.ConversationCommandResult(output)}) - - promptTokens, err := tokens.NewPromptTokenCounter(messages) - if err != nil { - return "", nil, trace.Wrap(err) - } - - resp, err := client.svc.CreateChatCompletion( - ctx, - openai.ChatCompletionRequest{ - Model: openai.GPT4, - Messages: messages, - }, - ) - - if err != nil { - return "", nil, trace.Wrap(err) - } - - completion := resp.Choices[0].Message.Content - completionTokens, err := tokens.NewSynchronousTokenCounter(completion) - - tc := &tokens.TokenCount{Prompt: tokens.TokenCounters{promptTokens}, Completion: tokens.TokenCounters{completionTokens}} - return completion, tc, trace.Wrap(err) -} - -// ClassifyMessage takes a user message, a list of categories, and uses the AI mode as a zero-shot classifier. -func (client *Client) ClassifyMessage(ctx context.Context, message string, classes map[string]string) (string, error) { - resp, err := client.svc.CreateChatCompletion( - ctx, - openai.ChatCompletionRequest{ - Model: openai.GPT4, - Messages: []openai.ChatCompletionMessage{ - {Role: openai.ChatMessageRoleSystem, Content: model.MessageClassificationPrompt(classes)}, - {Role: openai.ChatMessageRoleUser, Content: message}, - }, - }, - ) - - if err != nil { - return "", trace.Wrap(err) - } - - return resp.Choices[0].Message.Content, nil -} - -// ComputeEmbeddings takes a map of nodes and calls openAI to generate -// embeddings for those nodes. ComputeEmbeddings is responsible for -// implementing a retry mechanism if the embedding computation is flaky. -func (client *Client) ComputeEmbeddings(ctx context.Context, input []string) ([]embedding.Vector64, error) { - var results []embedding.Vector64 - for i := 0; maxOpenAIEmbeddingsPerRequest*i < len(input); i++ { - result, err := client.computeEmbeddings(ctx, paginateInput(input, i, maxOpenAIEmbeddingsPerRequest)) - if err != nil { - return nil, trace.Wrap(err) - } - for _, vector := range result { - results = append(results, embedding.Vector32to64(vector)) - } - } - return results, nil -} - -func paginateInput(input []string, page, pageSize int) []string { - begin := page * pageSize - var end int - if len(input) < (page+1)*pageSize { - end = len(input) - } else { - end = (page + 1) * pageSize - } - return input[begin:end] -} - -// computeEmbeddings calls the openAI embedding model with the provided input. -// This function should not be called directly, use ComputeEmbeddings instead -// to ensure input is properly batched. -func (client *Client) computeEmbeddings(ctx context.Context, input []string) ([]embedding.Vector32, error) { - req := openai.EmbeddingRequest{ - Input: input, - Model: openai.AdaEmbeddingV2, - } - - // Execute the query - resp, err := client.svc.CreateEmbeddings(ctx, req) - if err != nil { - return nil, trace.Wrap(err) - } - result := make([]embedding.Vector32, len(input)) - for i, item := range resp.Data { - result[i] = item.Embedding - } - return result, nil -} diff --git a/lib/ai/client_test.go b/lib/ai/client_test.go deleted file mode 100644 index 36b19a71b9c..00000000000 --- a/lib/ai/client_test.go +++ /dev/null @@ -1,112 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package ai - -import ( - "context" - "encoding/json" - "net/http/httptest" - "testing" - - "github.com/sashabaranov/go-openai" - "github.com/stretchr/testify/require" - "google.golang.org/grpc" - - assistpb "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" - "github.com/gravitational/teleport/lib/ai/model/output" - "github.com/gravitational/teleport/lib/ai/model/tools" - "github.com/gravitational/teleport/lib/ai/testutils" -) - -func TestRunTool_AuditQueryGeneration(t *testing.T) { - // Test setup: starting a mock openai server and creating the client - const generatedQuery = "SELECT user FROM session_start WHERE login='root'" - - responses := []string{ - // Then the audit query tool chooses to request session.start events - "session.start", - // Finally the tool builds a query based on the provided schemas - generatedQuery, - } - server := httptest.NewServer(testutils.GetTestHandlerFn(t, responses)) - t.Cleanup(server.Close) - - cfg := openai.DefaultConfig("secret-test-token") - cfg.BaseURL = server.URL - - client := NewClientFromConfig(cfg) - - // Doing the test: Check that the AuditQueryGeneration tool can be invoked - // through client.RunTool and validate its response. - ctx := context.Background() - toolCtx := &tools.ToolContext{User: "alice"} - response, _, err := client.RunTool(ctx, toolCtx, tools.AuditQueryGenerationToolName, "List users who connected to a server as root") - require.NoError(t, err) - message, ok := response.(*output.StreamingMessage) - require.True(t, ok) - require.Equal(t, generatedQuery, message.WaitAndConsume()) -} - -type mockEmbeddingGetter struct { - response []*assistpb.EmbeddedDocument -} - -func (m *mockEmbeddingGetter) GetAssistantEmbeddings(ctx context.Context, in *assistpb.GetAssistantEmbeddingsRequest, opts ...grpc.CallOption) (*assistpb.GetAssistantEmbeddingsResponse, error) { - return &assistpb.GetAssistantEmbeddingsResponse{Embeddings: m.response}, nil -} - -func TestRunTool_EmbeddingRetrieval(t *testing.T) { - // Test setup: starting a mock openai server and embedding getter, - // then create the client. - mock := &mockEmbeddingGetter{ - []*assistpb.EmbeddedDocument{ - { - Id: "1", - Content: "foo", - SimilarityScore: 1, - }, - { - Id: "2", - Content: "bar", - SimilarityScore: 0.9, - }, - }, - } - ctx := context.Background() - toolCtx := &tools.ToolContext{AssistEmbeddingServiceClient: mock} - - responses := make([]string, 0) - server := httptest.NewServer(testutils.GetTestHandlerFn(t, responses)) - t.Cleanup(server.Close) - - cfg := openai.DefaultConfig("secret-test-token") - cfg.BaseURL = server.URL - client := NewClientFromConfig(cfg) - - // Doing the test: Check that the EmbeddingRetrieval tool can be invoked - // through client.RunTool and validate its response. - input := tools.EmbeddingRetrievalToolInput{Question: "Find foobar"} - inputText, err := json.Marshal(input) - require.NoError(t, err) - response, _, err := client.RunTool(ctx, toolCtx, "Nodes names and labels retrieval", string(inputText)) - require.NoError(t, err) - message, ok := response.(*output.Message) - require.True(t, ok) - require.Equal(t, "foo\nbar\n", message.Content) -} diff --git a/lib/ai/embedding/embedding.go b/lib/ai/embedding/embedding.go deleted file mode 100644 index fc587fb394d..00000000000 --- a/lib/ai/embedding/embedding.go +++ /dev/null @@ -1,115 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package embedding - -import ( - "context" - "crypto/sha256" - - embeddingpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/embedding/v1" - "github.com/gravitational/teleport/lib/backend" -) - -// EmbeddingHash is the hash function that should be used to compute embedding -// hashes. -var EmbeddingHash = sha256.Sum256 - -// Sha256Hash is the hash of the embedded content. This hash allows to detect if -// the embedding is still up-to-date or if the content changed and the resource -// must be re-embedded. -type Sha256Hash = [sha256.Size]byte - -// Vector32 is an array of float64 that contains the result of the -// embedding process. OpenAI client returns []float32, hence Vector32 is the -// main type for handling vector data. -type Vector32 = []float32 - -// Vector64 is an array of float64 that contains the result of the embedding -// process. While OpenAI returns us 32-bit floats, the vector index uses methods -// requiring 64-bit floats. -type Vector64 = []float64 - -// Embedder is implemented for batch text embedding. Embedding can happen in -// place (with an embedding model, for example) or be done by a remote embedding -// service like OpenAI. -type Embedder interface { - // ComputeEmbeddings computes the embeddings of multiple strings. - // The embedding list follows the input order (e.g., result[i] is the - // embedding of input[i]). - ComputeEmbeddings(ctx context.Context, input []string) ([]Vector64, error) -} - -// Embedding contains a Teleport resource embedding. Embeddings are small semantic -// representations of larger and more complex data. Embeddings can be compared, -// the smaller the distance between two vectors, the closer the concepts are. -// Teleport Assist embeds resources to perform semantic search. -// The Embedding is named after the embedded resource id and kind. For example -// the SSH node "bastion-01" has the embedding "node/bastion-01". -type Embedding embeddingpb.Embedding - -// GetEmbeddedKind returns the kind of the resource that was embedded. -func (e *Embedding) GetEmbeddedKind() string { - return e.EmbeddedKind -} - -// GetName returns the Embedding name, composed of the embedded resource kind -// and the embedded resource ID. -func (e *Embedding) GetName() string { - return e.EmbeddedKind + string(backend.Separator) + e.EmbeddedId -} - -// GetEmbeddedID returns the ID of the resource that was embedded. -func (e *Embedding) GetEmbeddedID() string { - return e.EmbeddedId -} - -// GetVector returns the embedding vector -func (e *Embedding) GetVector() Vector64 { - return e.Vector -} - -// Dimensions returns the number of dimensions of the embedding -// Implements kdtree.Point interface -func (e *Embedding) Dimensions() int { - return len(e.Vector) -} - -// Dimension returns the value of the i-th dimension -// Implements kdtree.Point interface -func (e *Embedding) Dimension(i int) float64 { - return e.Vector[i] -} - -// NewEmbedding is an Embedding constructor. -func NewEmbedding(kind, id string, vector Vector64, hash Sha256Hash) *Embedding { - return &Embedding{ - EmbeddedKind: kind, - EmbeddedId: id, - EmbeddedHash: hash[:], - Vector: vector, - } -} - -func Vector32to64(vector32 Vector32) Vector64 { - vector64 := make(Vector64, len(vector32)) - for i, dimension := range vector32 { - vector64[i] = float64(dimension) - } - return vector64 -} diff --git a/lib/ai/embedding/serialization.go b/lib/ai/embedding/serialization.go deleted file mode 100644 index 2ba63bf3308..00000000000 --- a/lib/ai/embedding/serialization.go +++ /dev/null @@ -1,152 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package embedding - -import ( - "github.com/gravitational/trace" - "gopkg.in/yaml.v3" - - "github.com/gravitational/teleport/api/types" -) - -// SerializeNode converts a serializable resource into text ready to be fed to an -// embedding model. The YAML serialization function was chosen over JSON and -// CSV as it provided better results. -func SerializeResource(resource types.Resource) ([]byte, error) { - switch resource.GetKind() { - case types.KindNode: - return SerializeNode(resource.(types.Server)) - case types.KindKubernetesCluster: - return SerializeKubeCluster(resource.(types.KubeCluster)) - case types.KindApp: - return SerializeApp(resource.(types.Application)) - case types.KindDatabase: - return SerializeDatabase(resource.(types.Database)) - case types.KindWindowsDesktop: - return SerializeWindowsDesktop(resource.(types.WindowsDesktop)) - default: - return nil, trace.BadParameter("unknown resource kind %q", resource.GetKind()) - } -} - -// SerializeNode converts a type.Server into text ready to be fed to an -// embedding model. The YAML serialization function was chosen over JSON and -// CSV as it provided better results. -func SerializeNode(node types.Server) ([]byte, error) { - a := struct { - Name string `yaml:"name"` - Kind string `yaml:"kind"` - SubKind string `yaml:"subkind"` - Labels map[string]string `yaml:"labels"` - }{ - // Create artificial Name file for the node "name". Using node.GetName() as Name seems to confuse the model. - Name: node.GetHostname(), - Kind: types.KindNode, - SubKind: node.GetSubKind(), - Labels: node.GetAllLabels(), - } - text, err := yaml.Marshal(&a) - return text, trace.Wrap(err) -} - -// SerializeKubeCluster converts a type.KubeCluster into text ready to be fed to an -// embedding model. The YAML serialization function was chosen over JSON and -// CSV as it provided better results. -func SerializeKubeCluster(cluster types.KubeCluster) ([]byte, error) { - a := struct { - Name string `yaml:"name"` - Kind string `yaml:"kind"` - SubKind string `yaml:"subkind"` - Labels map[string]string `yaml:"labels"` - }{ - Name: cluster.GetName(), - Kind: types.KindKubernetesCluster, - SubKind: cluster.GetSubKind(), - Labels: cluster.GetAllLabels(), - } - text, err := yaml.Marshal(&a) - return text, trace.Wrap(err) -} - -// SerializeApp converts a type.Application into text ready to be fed to an -// embedding model. The YAML serialization function was chosen over JSON and -// CSV as it provided better results. -func SerializeApp(app types.Application) ([]byte, error) { - a := struct { - Name string `yaml:"name"` - Kind string `yaml:"kind"` - SubKind string `yaml:"subkind"` - Labels map[string]string `yaml:"labels"` - Description string `yaml:"description"` - }{ - Name: app.GetName(), - Kind: types.KindApp, - SubKind: app.GetSubKind(), - Labels: app.GetAllLabels(), - Description: app.GetDescription(), - } - text, err := yaml.Marshal(&a) - return text, trace.Wrap(err) -} - -// SerializeDatabase converts a type.Database into text ready to be fed to an -// embedding model. The YAML serialization function was chosen over JSON and -// CSV as it provided better results. -func SerializeDatabase(db types.Database) ([]byte, error) { - a := struct { - Name string `yaml:"name"` - Kind string `yaml:"kind"` - SubKind string `yaml:"subkind"` - Labels map[string]string `yaml:"labels"` - Type string `yaml:"type"` - Description string `yaml:"description"` - }{ - Name: db.GetName(), - Kind: types.KindDatabase, - SubKind: db.GetSubKind(), - Labels: db.GetAllLabels(), - Type: db.GetType(), - Description: db.GetDescription(), - } - text, err := yaml.Marshal(&a) - return text, trace.Wrap(err) -} - -// SerializeWindowsDesktop converts a type.WindowsDesktop into text ready to be fed to an -// embedding model. The YAML serialization function was chosen over JSON and -// CSV as it provided better results. -func SerializeWindowsDesktop(desktop types.WindowsDesktop) ([]byte, error) { - a := struct { - Name string `yaml:"name"` - Kind string `yaml:"kind"` - SubKind string `yaml:"subkind"` - Labels map[string]string `yaml:"labels"` - Address string `yaml:"address"` - ADDomain string `yaml:"ad_domain"` - }{ - Name: desktop.GetName(), - Kind: types.KindKubernetesCluster, - SubKind: desktop.GetSubKind(), - Labels: desktop.GetAllLabels(), - Address: desktop.GetAddr(), - ADDomain: desktop.GetDomain(), - } - text, err := yaml.Marshal(&a) - return text, trace.Wrap(err) -} diff --git a/lib/ai/embeddingprocessor.go b/lib/ai/embeddingprocessor.go deleted file mode 100644 index f01085e6b19..00000000000 --- a/lib/ai/embeddingprocessor.go +++ /dev/null @@ -1,331 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package ai - -import ( - "context" - "strings" - "time" - - "github.com/gravitational/trace" - "github.com/sirupsen/logrus" - "google.golang.org/protobuf/proto" - - embeddingpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/embedding/v1" - "github.com/gravitational/teleport/api/internalutils/stream" - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/api/utils/retryutils" - embeddinglib "github.com/gravitational/teleport/lib/ai/embedding" - "github.com/gravitational/teleport/lib/services" - streamutils "github.com/gravitational/teleport/lib/utils/stream" -) - -// maxEmbeddingAPISize is the maximum number of entities that can be embedded in a single API call. -const maxEmbeddingAPISize = 1000 - -// Embeddings implements the minimal interface used by the Embedding processor. -type Embeddings interface { - // GetAllEmbeddings returns all embeddings. - GetAllEmbeddings(ctx context.Context) stream.Stream[*embeddinglib.Embedding] - - // UpsertEmbedding creates or update a single ai.Embedding in the backend. - UpsertEmbedding(ctx context.Context, embedding *embeddinglib.Embedding) (*embeddinglib.Embedding, error) -} - -// MarshalEmbedding marshals the ai.Embedding resource to binary ProtoBuf. -func MarshalEmbedding(embedding *embeddinglib.Embedding) ([]byte, error) { - data, err := proto.Marshal((*embeddingpb.Embedding)(embedding)) - if err != nil { - return nil, trace.Wrap(err) - } - return data, nil -} - -// UnmarshalEmbedding unmarshals binary ProtoBuf into an ai.Embedding resource. -func UnmarshalEmbedding(bytes []byte) (*embeddinglib.Embedding, error) { - if len(bytes) == 0 { - return nil, trace.BadParameter("missing embedding data") - } - var embedding embeddingpb.Embedding - err := proto.Unmarshal(bytes, &embedding) - if err != nil { - return nil, trace.Wrap(err) - } - - return (*embeddinglib.Embedding)(&embedding), nil -} - -// EmbeddingHashMatches returns true if the hash of the embedding matches the -// given hash. -func EmbeddingHashMatches(embedding *embeddinglib.Embedding, hash embeddinglib.Sha256Hash) bool { - if len(embedding.EmbeddedHash) != 32 { - return false - } - - return *(*embeddinglib.Sha256Hash)(embedding.EmbeddedHash) == hash -} - -// BatchReducer is a helper that processes data in batches. -type BatchReducer[T, V any] struct { - data []T - batchSize int - processFn func(ctx context.Context, data []T) (V, error) -} - -// NewBatchReducer is a BatchReducer constructor. -func NewBatchReducer[T, V any](processFn func(ctx context.Context, data []T) (V, error), batchSize int) *BatchReducer[T, V] { - return &BatchReducer[T, V]{ - data: make([]T, 0), - batchSize: batchSize, - processFn: processFn, - } -} - -// Add adds a new item to the batch. If the batch is full, it will be processed -// and the result will be returned. Otherwise, a zero value will be returned. -// Finalize must be called to process the remaining data in the batch. -func (b *BatchReducer[T, V]) Add(ctx context.Context, data T) (V, error) { - b.data = append(b.data, data) - if len(b.data) >= b.batchSize { - val, err := b.processFn(ctx, b.data) - b.data = b.data[:0] - return val, trace.Wrap(err) - } - - var def V - return def, nil -} - -// Finalize processes the remaining data in the batch and returns the result. -func (b *BatchReducer[T, V]) Finalize(ctx context.Context) (V, error) { - if len(b.data) > 0 { - val, err := b.processFn(ctx, b.data) - b.data = b.data[:0] - return val, trace.Wrap(err) - } - - var def V - return def, nil -} - -// EmbeddingProcessorConfig is the configuration for EmbeddingProcessor. -type EmbeddingProcessorConfig struct { - AIClient embeddinglib.Embedder - EmbeddingSrv Embeddings - EmbeddingsRetriever *SimpleRetriever - NodeSrv *services.UnifiedResourceCache - Log logrus.FieldLogger - Jitter retryutils.Jitter -} - -// EmbeddingProcessor is responsible for processing nodes, generating embeddings -// and storing their embeddings in the backend. -type EmbeddingProcessor struct { - aiClient embeddinglib.Embedder - embeddingSrv Embeddings - embeddingsRetriever *SimpleRetriever - nodeSrv *services.UnifiedResourceCache - log logrus.FieldLogger - jitter retryutils.Jitter -} - -// NewEmbeddingProcessor returns a new EmbeddingProcessor. -func NewEmbeddingProcessor(cfg *EmbeddingProcessorConfig) *EmbeddingProcessor { - return &EmbeddingProcessor{ - aiClient: cfg.AIClient, - embeddingSrv: cfg.EmbeddingSrv, - embeddingsRetriever: cfg.EmbeddingsRetriever, - nodeSrv: cfg.NodeSrv, - log: cfg.Log, - jitter: cfg.Jitter, - } -} - -// resourceStringPair is a helper struct that pairs a resource with a data string. -type resourceStringPair struct { - resource types.Resource - data string -} - -// mapProcessFn is a helper function that maps a slice of resourceStringPair, -// compute embeddings and return them as a slice. -func (e *EmbeddingProcessor) mapProcessFn(ctx context.Context, data []*resourceStringPair) ([]*embeddinglib.Embedding, error) { - dataBatch := make([]string, 0, len(data)) - for _, pair := range data { - dataBatch = append(dataBatch, pair.data) - } - - embeddings, err := e.aiClient.ComputeEmbeddings(ctx, dataBatch) - if err != nil { - return nil, trace.Wrap(err) - } - - results := make([]*embeddinglib.Embedding, 0, len(embeddings)) - for i, embedding := range embeddings { - emb := embeddinglib.NewEmbedding(data[i].resource.GetKind(), - data[i].resource.GetName(), embedding, - embeddinglib.EmbeddingHash([]byte(data[i].data)), - ) - results = append(results, emb) - } - - return results, nil -} - -// Run runs the EmbeddingProcessor. -func (e *EmbeddingProcessor) Run(ctx context.Context, initialDelay, period time.Duration) error { - initTimer := time.NewTimer(initialDelay) - for { - select { - case <-ctx.Done(): - return ctx.Err() - case <-initTimer.C: - // Stop the timer after the initial delay. - initTimer.Stop() - e.process(ctx) - case <-time.After(e.jitter(period)): - e.process(ctx) - } - } -} - -// process updates embeddings for all resources once. -func (e *EmbeddingProcessor) process(ctx context.Context) { - batch := NewBatchReducer(e.mapProcessFn, - maxEmbeddingAPISize, // Max batch size allowed by OpenAI API, - ) - - e.log.Debugf("embedding processor started") - defer e.log.Debugf("embedding processor finished") - - embeddingsStream := e.embeddingSrv.GetAllEmbeddings(ctx) - unifiedResources, err := e.nodeSrv.GetUnifiedResources(ctx) - if err != nil { - e.log.Debugf("embedding processor failed with error: %v", err) - return - } - - resources := make([]types.Resource, len(unifiedResources)) - for i, unifiedResource := range unifiedResources { - resources[i] = unifiedResource - unifiedResources[i] = nil - } - - resourceStream := stream.Slice(resources) - - s := streamutils.NewZipStreams( - resourceStream, - embeddingsStream, - // On new resource callback. Add the resource to the batch. - func(resource types.Resource) error { - resourceData, err := embeddinglib.SerializeResource(resource) - if err != nil { - return trace.Wrap(err) - } - vectors, err := batch.Add(ctx, &resourceStringPair{resource, string(resourceData)}) - if err != nil { - return trace.Wrap(err) - } - if err := e.upsertEmbeddings(ctx, vectors); err != nil { - return trace.Wrap(err) - } - - return nil - }, - // On equal resource callback. Check if the resource's embedding hash matches - // the one in the backend. If not, add the resource to the batch. - func(resource types.Resource, embedding *embeddinglib.Embedding) error { - resourceData, err := embeddinglib.SerializeResource(resource) - if err != nil { - return trace.Wrap(err) - } - resourceHash := embeddinglib.EmbeddingHash(resourceData) - - if !EmbeddingHashMatches(embedding, resourceHash) { - vectors, err := batch.Add(ctx, &resourceStringPair{resource, string(resourceData)}) - if err != nil { - return trace.Wrap(err) - } - if err := e.upsertEmbeddings(ctx, vectors); err != nil { - return trace.Wrap(err) - } - } - return nil - }, - // On compare keys callback. Compare the keys for iteration. - func(resource types.Resource, embeddings *embeddinglib.Embedding) int { - return strings.Compare(resource.GetName(), embeddings.GetEmbeddedID()) - }, - ) - - if err := s.Process(); err != nil { - e.log.Warnf("Failed to generate nodes embedding: %v", err) - } - - // Process the remaining resources in the batch - vectors, err := batch.Finalize(ctx) - if err != nil { - e.log.Warnf("Failed to add node to batch: %v", err) - return - } - - if err := e.upsertEmbeddings(ctx, vectors); err != nil { - e.log.Warnf("Failed to upsert embeddings: %v", err) - - } - - if err := e.updateMemIndex(ctx); err != nil { - e.log.Warnf("Failed to update memory index: %v", err) - } -} - -// updateMemIndex is a helper function that updates the in-memory index with the -// latest embeddings. The new index is created and then swapped with the old one. -func (e *EmbeddingProcessor) updateMemIndex(ctx context.Context) error { - embeddingsIndex := NewSimpleRetriever() - embeddingsStream := e.embeddingSrv.GetAllEmbeddings(ctx) - - for embeddingsStream.Next() { - embedding := embeddingsStream.Item() - if !embeddingsIndex.Insert(embedding.GetEmbeddedID(), embedding) { - e.log.Warnf("Embeddings index is full, some resources can be missing") - break - } - } - - if err := embeddingsStream.Done(); err != nil { - return trace.Wrap(err) - } - - e.embeddingsRetriever.Swap(embeddingsIndex) - - return nil -} - -// upsertEmbeddings is a helper function that upserts the embeddings into the backend. -func (e *EmbeddingProcessor) upsertEmbeddings(ctx context.Context, rawEmbeddings []*embeddinglib.Embedding) error { - // Store the new embeddings into the backend - for _, embedding := range rawEmbeddings { - _, err := e.embeddingSrv.UpsertEmbedding(ctx, embedding) - if err != nil { - return trace.Wrap(err) - } - } - return nil -} diff --git a/lib/ai/embeddingprocessor_test.go b/lib/ai/embeddingprocessor_test.go deleted file mode 100644 index d53b99ed701..00000000000 --- a/lib/ai/embeddingprocessor_test.go +++ /dev/null @@ -1,313 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package ai_test - -import ( - "context" - "crypto/sha256" - "errors" - "fmt" - "strings" - "testing" - "time" - - "github.com/jonboulle/clockwork" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - - "github.com/gravitational/teleport/api/defaults" - "github.com/gravitational/teleport/api/internalutils/stream" - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/api/utils/retryutils" - "github.com/gravitational/teleport/lib/ai" - "github.com/gravitational/teleport/lib/ai/embedding" - "github.com/gravitational/teleport/lib/backend/memory" - "github.com/gravitational/teleport/lib/services" - "github.com/gravitational/teleport/lib/services/local" - "github.com/gravitational/teleport/lib/utils" -) - -func TestNodeEmbeddingGeneration(t *testing.T) { - t.Parallel() - - ctx, cancel := context.WithCancel(context.Background()) - t.Cleanup(cancel) - - clock := clockwork.NewFakeClock() - - // Test setup: crate a backend, presence service, the node watcher and - // the embeddings service. - bk, err := memory.New(memory.Config{ - Context: ctx, - Clock: clock, - }) - require.NoError(t, err) - - embedder := ai.MockEmbedder{ - TimesCalled: make(map[string]int), - } - events := local.NewEventsService(bk) - accessLists, err := local.NewAccessListService(bk, clock) - require.NoError(t, err) - resources := &mockResourceGetter{ - Presence: local.NewPresenceService(bk), - AccessLists: accessLists, - } - - cache, err := services.NewUnifiedResourceCache(ctx, services.UnifiedResourceCacheConfig{ - ResourceGetter: resources, - ResourceWatcherConfig: services.ResourceWatcherConfig{ - Component: "resource-watcher", - Client: events, - }, - }) - require.NoError(t, err) - - embeddings := local.NewEmbeddingsService(bk) - - processor := ai.NewEmbeddingProcessor(&ai.EmbeddingProcessorConfig{ - AIClient: &embedder, - EmbeddingSrv: embeddings, - EmbeddingsRetriever: ai.NewSimpleRetriever(), - NodeSrv: cache, - Log: utils.NewLoggerForTests(), - Jitter: retryutils.NewSeventhJitter(), - }) - - go func() { - err := processor.Run(ctx, 100*time.Millisecond, time.Second) - assert.ErrorIs(t, context.Canceled, err) - }() - - // Add some node servers. - const numInitialNodes = 5 - nodes := make([]types.Server, 0, numInitialNodes) - for i := 0; i < numInitialNodes; i++ { - node := makeNode(i + 1) - _, err = resources.UpsertNode(ctx, node) - require.NoError(t, err) - nodes = append(nodes, node) - } - - require.Eventually(t, func() bool { - items, err := stream.Collect(embeddings.GetAllEmbeddings(ctx)) - assert.NoError(t, err) - return len(items) == numInitialNodes - }, 14*time.Second, 200*time.Millisecond) - - nodesAcquired, err := resources.GetNodes(ctx, defaults.Namespace) - require.NoError(t, err) - - validateEmbeddings(t, - nodesAcquired, - embeddings.GetAllEmbeddings(ctx)) - - for k, v := range embedder.TimesCalled { - require.Equal(t, 1, v, "expected %v to be computed once, was %d", k, v) - } - - // Run once more and verify that only changed or newly inserted nodes get their embeddings calculated - node1 := nodes[0] - node1.GetMetadata().Labels["foo"] = "bar" - _, err = resources.UpsertNode(ctx, node1) - require.NoError(t, err) - node6 := makeNode(6) - _, err = resources.UpsertNode(ctx, node6) - require.NoError(t, err) - - // Since nodes are streamed in ascending order by names, when embeddings for node6 are calculated, - // we can be sure that our recent changes have been fully processed - require.Eventually(t, func() bool { - items, err := stream.Collect(embeddings.GetAllEmbeddings(ctx)) - assert.NoError(t, err) - return len(items) == numInitialNodes+1 - }, 7*time.Second, 200*time.Millisecond) - - for k, v := range embedder.TimesCalled { - expected := 1 - if strings.Contains(k, "node1") { - expected = 2 - } - require.Equal(t, expected, v, "expected embedding for %q to be computed %d times, got computed %d times", k, expected, v) - } - - nodesAcquired, err = resources.GetNodes(ctx, defaults.Namespace) - require.NoError(t, err) - - validateEmbeddings(t, - nodesAcquired, - embeddings.GetAllEmbeddings(ctx)) -} - -func TestMarshallUnmarshallEmbedding(t *testing.T) { - // We test that float precision is above six digits - initial := embedding.NewEmbedding(types.KindNode, "foo", embedding.Vector64{0.1234567, 1, 1}, sha256.Sum256([]byte("test"))) - - marshaled, err := ai.MarshalEmbedding(initial) - require.NoError(t, err) - - final, err := ai.UnmarshalEmbedding(marshaled) - require.NoError(t, err) - - require.Equal(t, initial.EmbeddedId, final.EmbeddedId) - require.Equal(t, initial.EmbeddedKind, final.EmbeddedKind) - require.Equal(t, initial.EmbeddedHash, final.EmbeddedHash) - require.Equal(t, initial.Vector, final.Vector) -} - -func makeNode(num int) types.Server { - node, _ := types.NewServer(fmt.Sprintf("node%d", num), types.KindNode, types.ServerSpecV2{ - Addr: "127.0.0.1:1234", - Hostname: fmt.Sprintf("node%d", num), - CmdLabels: map[string]types.CommandLabelV2{ - "version": {Result: "v8"}, - "hostname": {Result: fmt.Sprintf("node%d.example.com", num)}, - }, - }) - return node -} - -func validateEmbeddings(t *testing.T, nodes []types.Server, embeddingsStream stream.Stream[*embedding.Embedding]) { - t.Helper() - - embeddings, err := stream.Collect(embeddingsStream) - require.NoError(t, err) - - require.Equal(t, len(nodes), len(embeddings), "Number of nodes and embeddings should be equal") - - for i, node := range nodes { - emb := embeddings[i] - - require.Equal(t, node.GetName(), emb.GetEmbeddedID(), "Node ID and embedding ID should be equal") - require.Equal(t, types.KindNode, emb.GetEmbeddedKind(), "Node kind and embedding kind should be equal") - } -} - -func Test_batchReducer_Add(t *testing.T) { - t.Parallel() - - // Sum process function - used for simplicity - sumFn := func(ctx context.Context, data []int) (int, error) { - sum := 0 - for _, d := range data { - sum += d - } - return sum, nil - } - - type testCase struct { - // Test case name - name string - // Process batch size - batchSize int - // Input data - data []int - // Function to process batch - processFn func(ctx context.Context, data []int) (int, error) - // Expected result on Add - want []int - // Expected result on Finalize - finalizeResult int - // Expected error - wantErr assert.ErrorAssertionFunc - } - - tests := []testCase{ - { - name: "empty", - batchSize: 100, - data: []int{}, - want: []int{}, - finalizeResult: 0, - processFn: sumFn, - wantErr: assert.NoError, - }, - { - name: "one element", - batchSize: 100, - data: []int{1}, - want: []int{0}, - finalizeResult: 1, - processFn: sumFn, - wantErr: assert.NoError, - }, - { - name: "many elements", - batchSize: 3, - data: []int{1, 1, 1, 1}, - want: []int{0, 0, 3, 0}, - finalizeResult: 1, - processFn: sumFn, - wantErr: assert.NoError, - }, - { - name: "propagate error", - batchSize: 2, - data: []int{0}, - want: []int{0}, - processFn: func(ctx context.Context, data []int) (int, error) { - return 0, errors.New("error") - }, - wantErr: assert.Error, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - ctx := context.Background() - br := ai.NewBatchReducer[int, int](tt.processFn, tt.batchSize) - - for i, d := range tt.data { - got, err := br.Add(ctx, d) - require.NoError(t, err) - assert.Equalf(t, tt.want[i], got, "Add(%v)", tt.data) - } - - got, err := br.Finalize(ctx) - if !tt.wantErr(t, err, fmt.Sprintf("Finalize(%v)", tt.data)) { - return - } - assert.Equalf(t, tt.finalizeResult, got, "Finalize(%v)", tt.data) - }) - } -} - -type mockResourceGetter struct { - services.Presence - services.AccessLists -} - -func (m *mockResourceGetter) GetDatabaseServers(_ context.Context, _ string, _ ...services.MarshalOption) ([]types.DatabaseServer, error) { - return nil, nil -} - -func (m *mockResourceGetter) GetKubernetesServers(_ context.Context) ([]types.KubeServer, error) { - return nil, nil -} - -func (m *mockResourceGetter) GetApplicationServers(_ context.Context, _ string) ([]types.AppServer, error) { - return nil, nil -} - -func (m *mockResourceGetter) GetWindowsDesktops(_ context.Context, _ types.WindowsDesktopFilter) ([]types.WindowsDesktop, error) { - return nil, nil -} - -func (m *mockResourceGetter) ListSAMLIdPServiceProviders(_ context.Context, _ int, _ string) ([]types.SAMLIdPServiceProvider, string, error) { - return nil, "", nil -} diff --git a/lib/ai/mock_embedder.go b/lib/ai/mock_embedder.go deleted file mode 100644 index ab1bc600f46..00000000000 --- a/lib/ai/mock_embedder.go +++ /dev/null @@ -1,54 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package ai - -import ( - "context" - "crypto/sha256" - "strings" - "sync" - - "github.com/gravitational/teleport/lib/ai/embedding" -) - -// MockEmbedder returns embeddings based on the sha256 hash function. Those -// embeddings have no semantic meaning but ensure different embedded content -// provides different embeddings. -type MockEmbedder struct { - mu sync.Mutex - TimesCalled map[string]int -} - -func (m *MockEmbedder) ComputeEmbeddings(_ context.Context, input []string) ([]embedding.Vector64, error) { - m.mu.Lock() - defer m.mu.Unlock() - - result := make([]embedding.Vector64, len(input)) - for i, text := range input { - name := strings.Split(text, "\n")[0] - m.TimesCalled[name]++ - hash := sha256.Sum256([]byte(text)) - vector := make(embedding.Vector64, len(hash)) - for j, x := range hash { - vector[j] = 1 / float64(int(x)+1) - } - result[i] = vector - } - return result, nil -} diff --git a/lib/ai/model/agent.go b/lib/ai/model/agent.go deleted file mode 100644 index de156a23387..00000000000 --- a/lib/ai/model/agent.go +++ /dev/null @@ -1,436 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package model - -import ( - "context" - "encoding/json" - "fmt" - "strings" - "time" - - "github.com/gravitational/trace" - "github.com/sashabaranov/go-openai" - log "github.com/sirupsen/logrus" - - "github.com/gravitational/teleport/lib/ai/model/output" - "github.com/gravitational/teleport/lib/ai/model/tools" - "github.com/gravitational/teleport/lib/ai/tokens" -) - -const ( - // The internal name used to create actions when the agent encounters an error, such as when parsing output. - actionException = "_Exception" - - // The maximum amount of thought <-> observation iterations the agent is allowed to perform. - maxIterations = 15 - - // The maximum amount of time the agent is allowed to spend before yielding a final answer. - maxElapsedTime = 5 * time.Minute - - // The special header the LLM has to respond with to indicate it's done. - finalResponseHeader = "" -) - -// NewAgent creates a new agent. The Assist agent which defines the model responsible for the Assist feature. -func NewAgent(toolCtx *tools.ToolContext, tools ...tools.Tool) *Agent { - return &Agent{tools, toolCtx} -} - -// Agent is a model storing static state which defines some properties of the chat model. -type Agent struct { - tools []tools.Tool - toolCtx *tools.ToolContext -} - -// AgentAction is an event type representing the decision to take a single action, typically a tool invocation. -type AgentAction struct { - // The action to take, typically a tool name. - Action string `json:"action"` - - // The input to the action, varies depending on the action. - Input string `json:"input"` - - // The log is either a direct tool response or a thought prompt correlated to the input. - Log string `json:"log"` - - // The reasoning is a string describing the reasoning behind the action. - Reasoning string `json:"reasoning"` -} - -// agentFinish is an event type representing the decision to finish a thought -// loop and return a final text answer to the user. -type agentFinish struct { - // output must be Message, StreamingMessage, CompletionCommand, AccessRequest. - output any -} - -type executionState struct { - llm *openai.Client - chatHistory []openai.ChatCompletionMessage - humanMessage openai.ChatCompletionMessage - intermediateSteps []AgentAction - observations []string - tokenCount *tokens.TokenCount -} - -// PlanAndExecute runs the agent with a given input until it arrives at a text answer it is satisfied -// with or until it times out. -func (a *Agent) PlanAndExecute(ctx context.Context, llm *openai.Client, chatHistory []openai.ChatCompletionMessage, humanMessage openai.ChatCompletionMessage, progressUpdates func(*AgentAction)) (any, *tokens.TokenCount, error) { - log.Trace("entering agent think loop") - iterations := 0 - start := time.Now() - tookTooLong := func() bool { return iterations > maxIterations || time.Since(start) > maxElapsedTime } - state := &executionState{ - llm: llm, - chatHistory: chatHistory, - humanMessage: humanMessage, - intermediateSteps: make([]AgentAction, 0), - observations: make([]string, 0), - tokenCount: tokens.NewTokenCount(), - } - - for { - log.Tracef("performing iteration %v of loop, %v seconds elapsed", iterations, int(time.Since(start).Seconds())) - - // This is intentionally not context-based, as we want to finish the current step before exiting - // and the concern is not that we're stuck but that we're taking too long over multiple iterations. - if tookTooLong() { - return nil, nil, trace.Errorf("timeout: agent took too long to finish") - } - - output, err := a.takeNextStep(ctx, state, progressUpdates) - if err != nil { - return nil, nil, trace.Wrap(err) - } - - if output.finish != nil { - log.Tracef("agent finished with output: %#v", output.finish.output) - - return output.finish.output, state.tokenCount, nil - } - - if output.action != nil { - state.intermediateSteps = append(state.intermediateSteps, *output.action) - state.observations = append(state.observations, output.observation) - } - - iterations++ - } -} - -func (a *Agent) DoAction(ctx context.Context, llm *openai.Client, action *AgentAction) (any, *tokens.TokenCount, error) { - state := &executionState{ - llm: llm, - tokenCount: tokens.NewTokenCount(), - } - out, err := a.doAction(ctx, state, action) - if err != nil { - return nil, nil, trace.Wrap(err) - } - switch { - case out.finish != nil: - // If the tool already breaks execution, we don't have to do anything - return out.finish.output, state.tokenCount, nil - case out.observation != "": - // If the tool doesn't break execution and returns a single observation, - // we wrap the observation in a Message. - return &output.Message{Content: out.observation}, state.tokenCount, nil - default: - return nil, state.tokenCount, trace.Errorf("action %s did not end execution nor returned an observation", action.Action) - } -} - -// stepOutput represents the inputs and outputs of a single thought step. -type stepOutput struct { - // if the agent is done, finish is set. - finish *agentFinish - - // if the agent is not done, action is set together with observation. - action *AgentAction - observation string -} - -func (e *executionState) isRepeatAction(action *AgentAction) bool { - for _, previousAction := range e.intermediateSteps { - if previousAction.Action == action.Action && previousAction.Input == action.Input { - return true - } - } - - return false -} - -func (a *Agent) takeNextStep(ctx context.Context, state *executionState, progressUpdates func(*AgentAction)) (stepOutput, error) { - log.Trace("agent entering takeNextStep") - defer log.Trace("agent exiting takeNextStep") - - action, finish, err := a.plan(ctx, state) - if output.IsInvalidOutputError(err) { - log.Tracef("agent encountered an invalid output error: %v, attempting to recover", err) - action := &AgentAction{ - Action: actionException, - Log: "Invalid or incomplete response: " + err.Error(), - } - - // The exception tool is currently a bit special, the observation is always equal to the input. - // We can expand on this in the future to make it handle errors better. - log.Tracef("agent decided on action %v and received observation %v", action.Action, action.Input) - return stepOutput{action: action, observation: action.Log}, nil - } - if err != nil { - log.Tracef("agent encountered an error: %v", err) - return stepOutput{}, trace.Wrap(err) - } - - // If finish is set, the agent is done and did not call upon any tool. - if finish != nil { - log.Trace("agent picked finish, returning") - return stepOutput{finish: finish}, nil - } - - // we check against repeat actions to get the LLM out of confusion loops faster. - if state.isRepeatAction(action) { - return stepOutput{action: action, observation: "You've already ran this tool with this input."}, nil - } - - // If action is set, the agent is not done and called upon a tool. - progressUpdates(action) - - return a.doAction(ctx, state, action) -} - -func (a *Agent) doAction(ctx context.Context, state *executionState, action *AgentAction) (stepOutput, error) { - var tool tools.Tool - for _, candidate := range a.tools { - if candidate.Name() == action.Action { - tool = candidate - break - } - } - - if tool == nil { - log.Tracef("agent picked an unknown tool %v", action.Action) - action := &AgentAction{ - Action: actionException, - Log: fmt.Sprintf("No tool with name %s exists.", action.Action), - } - - return stepOutput{action: action, observation: action.Log}, nil - } - - // Here we switch on the tool type because even though all tools are presented as equal to the LLM - // some are marked special and break the typical tool execution loop; those are handled here instead. - switch tool := tool.(type) { - case *tools.CommandExecutionTool: - completion, err := tool.ParseInput(action.Input) - if err != nil { - action := &AgentAction{ - Action: actionException, - Log: "Invalid or incomplete response: " + err.Error(), - } - - return stepOutput{action: action, observation: action.Log}, nil - } - - log.Tracef("agent decided on command execution, let's translate to an agentFinish") - return stepOutput{finish: &agentFinish{output: completion}}, nil - case *tools.AccessRequestCreateTool: - accessRequest, err := tool.ParseInput(action.Input) - if err != nil { - action := &AgentAction{ - Action: actionException, - Log: "Invalid or incomplete response: " + err.Error(), - } - - return stepOutput{action: action, observation: action.Log}, nil - } - - return stepOutput{finish: &agentFinish{output: accessRequest}}, nil - case *tools.CommandGenerationTool: - input, err := tool.ParseInput(action.Input) - if err != nil { - action := &AgentAction{ - Action: actionException, - Log: "Invalid or incomplete response: " + err.Error(), - } - - return stepOutput{action: action, observation: action.Log}, nil - } - completion := &output.GeneratedCommand{ - Command: input.Command, - } - - log.Tracef("agent decided on command generation, let's translate to an agentFinish") - return stepOutput{finish: &agentFinish{output: completion}}, nil - case *tools.AuditQueryGenerationTool: - log.Tracef("Tool called with input:'%s'", action.Input) - tableName, err := tool.ChooseEventTable(ctx, action.Input, state.tokenCount) - // If the query was not answerable by audit logs, - // we return to the agent thinking loop and tell that the tool failed - if trace.IsNotFound(err) { - return stepOutput{action: action, observation: err.Error()}, nil - } - if err != nil { - return stepOutput{}, trace.Wrap(err) - } - - log.Tracef("Tool chose to query table '%s'", tableName) - response, err := tool.GenerateQuery(ctx, tableName, action.Input, state.tokenCount) - if err != nil { - return stepOutput{}, trace.Wrap(err) - } - - return stepOutput{finish: &agentFinish{output: response}}, nil - default: - runOut, err := tool.Run(ctx, a.toolCtx, action.Input) - if err != nil { - return stepOutput{}, trace.Wrap(err) - } - return stepOutput{action: action, observation: runOut}, nil - } -} - -func (a *Agent) plan(ctx context.Context, state *executionState) (*AgentAction, *agentFinish, error) { - scratchpad := a.constructScratchpad(state.intermediateSteps, state.observations) - prompt := a.createPrompt(state.chatHistory, scratchpad, state.humanMessage) - promptTokenCount, err := tokens.NewPromptTokenCounter(prompt) - if err != nil { - return nil, nil, trace.Wrap(err) - } - state.tokenCount.AddPromptCounter(promptTokenCount) - - stream, err := state.llm.CreateChatCompletionStream( - ctx, - openai.ChatCompletionRequest{ - Model: openai.GPT432K, - Messages: prompt, - Temperature: 0.3, - Stream: true, - }, - ) - if err != nil { - return nil, nil, trace.Wrap(err) - } - - deltas := output.StreamToDeltas(stream) - - action, finish, completionTokenCounter, err := parsePlanningOutput(deltas) - state.tokenCount.AddCompletionCounter(completionTokenCounter) - return action, finish, trace.Wrap(err) -} - -func (a *Agent) createPrompt(chatHistory, agentScratchpad []openai.ChatCompletionMessage, humanMessage openai.ChatCompletionMessage) []openai.ChatCompletionMessage { - prompt := make([]openai.ChatCompletionMessage, 0) - prompt = append(prompt, chatHistory...) - toolList := strings.Builder{} - toolNames := make([]string, 0, len(a.tools)) - for _, tool := range a.tools { - toolNames = append(toolNames, tool.Name()) - toolList.WriteString("> ") - toolList.WriteString(tool.Name()) - toolList.WriteString(": ") - toolList.WriteString(tool.Description()) - toolList.WriteString("\n") - } - - if len(a.tools) == 0 { - toolList.WriteString("No tools available.") - } - - formatInstructions := conversationParserFormatInstructionsPrompt(toolNames) - newHumanMessage := conversationToolUsePrompt(toolList.String(), formatInstructions, humanMessage.Content) - prompt = append(prompt, openai.ChatCompletionMessage{ - Role: openai.ChatMessageRoleUser, - Content: newHumanMessage, - }) - - prompt = append(prompt, agentScratchpad...) - return prompt -} - -func (a *Agent) constructScratchpad(intermediateSteps []AgentAction, observations []string) []openai.ChatCompletionMessage { - var thoughts []openai.ChatCompletionMessage - for i, action := range intermediateSteps { - if len(action.Reasoning) != 0 { - thoughts = append(thoughts, openai.ChatCompletionMessage{ - Role: openai.ChatMessageRoleAssistant, - Content: action.Reasoning, - }) - } - - thoughts = append(thoughts, openai.ChatCompletionMessage{ - Role: openai.ChatMessageRoleUser, - Content: conversationToolResponse(observations[i]), - }) - } - - return thoughts -} - -// PlanOutput describes the expected JSON output after asking it to plan its next action. -type PlanOutput struct { - Action string `json:"action"` - ActionInput any `json:"action_input"` - Reasoning string `json:"reasoning"` -} - -// parsePlanningOutput parses the output of the model after asking it to plan its next action -// and returns the appropriate event type or an error. -func parsePlanningOutput(deltas <-chan string) (*AgentAction, *agentFinish, tokens.TokenCounter, error) { - var text string - for delta := range deltas { - text += delta - - if strings.HasPrefix(text, finalResponseHeader) { - message, tc, err := output.NewStreamingMessage(deltas, text, finalResponseHeader) - if err != nil { - return nil, nil, nil, trace.Wrap(err) - } - return nil, &agentFinish{output: message}, tc, nil - } - } - - completionTokenCount, err := tokens.NewSynchronousTokenCounter(text) - if err != nil { - return nil, nil, nil, trace.Wrap(err) - } - - log.Tracef("received planning output: \"%v\"", text) - if outputString, found := strings.CutPrefix(text, finalResponseHeader); found { - return nil, &agentFinish{output: &output.Message{Content: outputString}}, completionTokenCount, nil - } - - response, err := output.ParseJSONFromModel[PlanOutput](text) - if err != nil { - log.WithError(err).Trace("failed to parse planning output") - return nil, nil, nil, trace.Wrap(err) - } - - if v, ok := response.ActionInput.(string); ok { - return &AgentAction{Action: response.Action, Input: v}, nil, completionTokenCount, nil - } else { - input, err := json.Marshal(response.ActionInput) - if err != nil { - return nil, nil, nil, trace.Wrap(err) - } - - return &AgentAction{Action: response.Action, Input: string(input), Reasoning: response.Reasoning}, nil, completionTokenCount, nil - } -} diff --git a/lib/ai/model/output/error.go b/lib/ai/model/output/error.go deleted file mode 100644 index e11abf20a5b..00000000000 --- a/lib/ai/model/output/error.go +++ /dev/null @@ -1,60 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package output - -import ( - "errors" - "fmt" - - "github.com/gravitational/trace" -) - -// NewInvalidOutputError builds an error caused by the output of an LLM. -func NewInvalidOutputError(coarse, detail string) error { - return &invalidOutputError{ - coarse: coarse, - detail: detail, - } -} - -// IsInvalidOutputError returns true if the error is an invalidOutputError. -func IsInvalidOutputError(err error) bool { - var invalidOutputError *invalidOutputError - return errors.As(trace.Unwrap(err), &invalidOutputError) -} - -// invalidOutputError represents an error caused by the output of an LLM. -// These may be used automatically by the agent loop to attempt to correct an output until it is valid. -type invalidOutputError struct { - coarse string - detail string -} - -// newInvalidOutputErrorWithParseError creates a new invalidOutputError assuming a JSON parse error. -func newInvalidOutputErrorWithParseError(err error) *invalidOutputError { - return &invalidOutputError{ - coarse: "json parse error", - detail: err.Error(), - } -} - -// Error returns a string representation of the error. This is used to satisfy the error interface. -func (o *invalidOutputError) Error() string { - return fmt.Sprintf("%v: %v", o.coarse, o.detail) -} diff --git a/lib/ai/model/output/messages.go b/lib/ai/model/output/messages.go deleted file mode 100644 index a43873ddba8..00000000000 --- a/lib/ai/model/output/messages.go +++ /dev/null @@ -1,79 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package output - -import "strings" - -// Message represents a new message within a live conversation. -type Message struct { - Content string -} - -// StreamingMessage represents a new message that is being streamed from the LLM. -type StreamingMessage struct { - Parts <-chan string -} - -// WaitAndConsume waits until the message stream is over and returns the full message. -// This can only be called once on a message as it empties its Parts channel. -func (msg *StreamingMessage) WaitAndConsume() string { - sb := strings.Builder{} - for part := range msg.Parts { - sb.WriteString(part) - } - return sb.String() -} - -// Label represents a label returned by OpenAI's completion API. -type Label struct { - Key string `json:"key"` - Value string `json:"value"` -} - -// CompletionCommand represents a command suggestion returned by OpenAI's completion API. -type CompletionCommand struct { - Command string `json:"command,omitempty"` - Nodes []string `json:"nodes,omitempty"` - Labels []Label `json:"labels,omitempty"` -} - -// GeneratedCommand represents a Bash command generated by LLM. -type GeneratedCommand struct { - Command string `json:"command"` -} - -// AccessRequest represents an access request suggestion returned by OpenAI's completion API. -type AccessRequest struct { - Roles []string `json:"roles"` - Resources []Resource `json:"resources"` - Reason string `json:"reason"` - SuggestedReviewers []string `json:"suggested_reviewers"` -} - -// Resource represents a resource suggestion returned by OpenAI's completion API. -type Resource struct { - // The resource type. - Type string `json:"type"` - - // The resource name. - Name string `json:"id"` - - // Set if a display-friendly alternative name is available. - FriendlyName string `json:"friendlyName,omitempty"` -} diff --git a/lib/ai/model/output/parsejson.go b/lib/ai/model/output/parsejson.go deleted file mode 100644 index 7bedaa26927..00000000000 --- a/lib/ai/model/output/parsejson.go +++ /dev/null @@ -1,48 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package output - -import ( - "encoding/json" - "strings" -) - -// ParseJSONFromModel parses a JSON object from the model output and attempts to sanitize contaminant text -// to avoid triggering self-correction due to some natural language being bundled with the JSON. -// The output type is generic, and thus the structure of the expected JSON varies depending on T. -func ParseJSONFromModel[T any](text string) (T, error) { - cleaned := strings.TrimSpace(text) - if strings.Contains(cleaned, "```json") { - cleaned = strings.Split(cleaned, "```json")[1] - } - if strings.Contains(cleaned, "```") { - cleaned = strings.Split(cleaned, "```")[0] - } - cleaned = strings.TrimPrefix(cleaned, "```json") - cleaned = strings.TrimPrefix(cleaned, "```") - cleaned = strings.TrimSuffix(cleaned, "```") - cleaned = strings.TrimSpace(cleaned) - var output T - err := json.Unmarshal([]byte(cleaned), &output) - if err != nil { - return output, newInvalidOutputErrorWithParseError(err) - } - - return output, nil -} diff --git a/lib/ai/model/output/streaming.go b/lib/ai/model/output/streaming.go deleted file mode 100644 index 59a96ecae88..00000000000 --- a/lib/ai/model/output/streaming.go +++ /dev/null @@ -1,83 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package output - -import ( - "errors" - "io" - "strings" - - "github.com/gravitational/trace" - "github.com/sashabaranov/go-openai" - log "github.com/sirupsen/logrus" - - "github.com/gravitational/teleport/lib/ai/tokens" -) - -// StreamToDeltas converts an openai.CompletionStream into a channel of strings. -// This channel can then be consumed manually to search for specific markers, -// or directly converted into a StreamingMessage with NewStreamingMessage. -func StreamToDeltas(stream *openai.ChatCompletionStream) chan string { - deltas := make(chan string) - go func() { - defer close(deltas) - - for { - response, err := stream.Recv() - if errors.Is(err, io.EOF) { - return - } else if err != nil { - log.Tracef("agent encountered an error while streaming: %v", err) - return - } - - delta := response.Choices[0].Delta.Content - deltas <- delta - } - }() - return deltas -} - -// NewStreamingMessage takes a string channel and converts it to -// a StreamingMessage. -// If content was already streamed, it must be passed through the alreadyStreamed parameter. -// If the already streamed content contains a prefix that must be stripped -// (like a marker to identify the kind of response the model is providing), -// the prefix can be passed through the prefix parameter. It will be stripped -// but will still be reflected in the token count. -func NewStreamingMessage(deltas <-chan string, alreadyStreamed, prefix string) (*StreamingMessage, *tokens.AsynchronousTokenCounter, error) { - parts := make(chan string) - streamingTokenCounter, err := tokens.NewAsynchronousTokenCounter(alreadyStreamed) - if err != nil { - return nil, nil, trace.Wrap(err) - } - go func() { - defer close(parts) - - parts <- strings.TrimPrefix(alreadyStreamed, prefix) - for delta := range deltas { - parts <- delta - errCount := streamingTokenCounter.Add() - if errCount != nil { - log.WithError(errCount).Debug("Failed to add streamed completion text to the token counter") - } - } - }() - return &StreamingMessage{Parts: parts}, streamingTokenCounter, nil -} diff --git a/lib/ai/model/prompt.go b/lib/ai/model/prompt.go deleted file mode 100644 index 7b76e48caab..00000000000 --- a/lib/ai/model/prompt.go +++ /dev/null @@ -1,131 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package model - -import ( - "fmt" - "strings" -) - -const PromptSummarizeTitle = `You will be given a message. Create a short summary of that message. -Respond only with summary, nothing else.` - -const PromptSummarizeCommand = `You will be given a chat history and a command output. Based on the history context, extract relevant information from the command output and write a short summary of the command output. -Respond only with summary, nothing else.` - -const InitialAIResponse = `Hey, I'm Teleport - a powerful tool that can assist you in managing your Teleport cluster via OpenAI GPT-4.` - -func PromptCharacter(username string) string { - return fmt.Sprintf(`You are Teleport, a tool that users can use to connect to Linux servers and run relevant commands, as well as have a conversation. -A Teleport cluster is a connectivity layer that allows access to a set of servers. Servers may also be referred to as nodes. -Nodes sometimes have labels such as "production" and "staging" assigned to them. Labels are used to group nodes together. -You will engage in professional conversation with the user and help accomplish tasks such as executing tasks -within the cluster or answering relevant questions about Teleport, Linux or the cluster itself. - -You possess advanced capabilities to think and reason in multiple steps and use the available tools to accomplish the task at hand in a way a human would expect you to. - -You are not permitted to engage in conversation that is not related to Teleport, Linux or the cluster itself. -If this user asks such an unrelated question, you must concisely respond that it is beyond your scope of knowledge. - -You are talking to %v.`, username) -} - -func conversationParserFormatInstructionsPrompt(toolnames []string) string { - return fmt.Sprintf(`RESPONSE FORMAT INSTRUCTIONS ----------------------------- - -When responding to me, please output a response in one of two formats: - -**Option 1:** -Use this if you want the human to use a tool. -Markdown code snippet formatted in the following schema: - -%vjson -{ - "action": string \\ The action to take. Must be one of %v - "action_input": string \\ The input to the action - "reasoning": string \\ Your reasoning for taking this action -} -%v - -**Option #2:** -Use this if you want to respond directly to the human or you want to ask the human a question to gather more information. -You should avoid asking too many questions when you have other options available to you as it may be perceived as annoying. -But asking is far better than guessing or making assumptions. -Text with the hardcoded header %v followed by your response as below: - -%v -YOUR RESPONSE HERE`, "```", toolnames, "```", finalResponseHeader, finalResponseHeader, - ) -} - -func conversationToolUsePrompt(tools string, formatInstructions string, userInput string) string { - return fmt.Sprintf(`TOOLS ------- -Assistant can ask the user to use tools to look up information that may be helpful in answering the users original question. The tools the human can use are: - -%v - -%v - -USER'S INPUT --------------------- -Here is the user's input (remember to respond with a markdown code snippet of a json blob with a single action, and NOTHING else): - -%v`, tools, formatInstructions, userInput) -} - -func conversationToolResponse(toolResponse string) string { - return fmt.Sprintf(`TOOL RESPONSE: ---------------------- - -%v - -USER'S INPUT --------------------- - -Okay, so what is the response to my last comment? If using information obtained from the tools you must mention it explicitly without mentioning the tool names - I have forgotten all TOOL RESPONSES! Remember to respond with a markdown code snippet of a json blob with a single action, and NOTHING else.`, toolResponse) -} - -func ConversationCommandResult(result map[string][]byte) string { - var message strings.Builder - for node, output := range result { - message.WriteString(fmt.Sprintf(`Command ran on node "%s" and produced the following output:\n`, node)) - message.WriteString(string(output)) - message.WriteString("\n") - } - message.WriteString("Based on the chat history, extract relevant information out of the command output and write a summary. " + - "For error messages suggest a solution if possible. The solution can contain a Linux command or a description.") - return message.String() -} - -func MessageClassificationPrompt(classes map[string]string) string { - var classList strings.Builder - for name, description := range classes { - classList.WriteString(fmt.Sprintf("- `%s` (%s)\n", name, description)) - } - - return fmt.Sprintf(`Teleport is a tool that provides access to servers, kubernetes clusters, databases, and applications. All connected Teleport resources are called a cluster. Server resources might be called nodes. - -Classify the provided message between the following categories: - -%v - -Answer only with the category name. Nothing else.`, classList.String()) -} diff --git a/lib/ai/model/tools/accessrequest.go b/lib/ai/model/tools/accessrequest.go deleted file mode 100644 index 2fbe64d4811..00000000000 --- a/lib/ai/model/tools/accessrequest.go +++ /dev/null @@ -1,307 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package tools - -import ( - "context" - "fmt" - "strconv" - "strings" - "sync" - "time" - - "github.com/google/uuid" - "github.com/gravitational/trace" - "golang.org/x/sync/errgroup" - "gopkg.in/yaml.v3" - - "github.com/gravitational/teleport/api/client/proto" - "github.com/gravitational/teleport/api/types" - modeloutput "github.com/gravitational/teleport/lib/ai/model/output" -) - -type AccessRequestListRequestableRolesTool struct{} - -func (*AccessRequestListRequestableRolesTool) Name() string { - return "List Requestable Roles" -} - -func (*AccessRequestListRequestableRolesTool) Description() string { - return "List all roles that can be requested via access requests." -} - -func (a *AccessRequestListRequestableRolesTool) Run(ctx context.Context, toolCtx *ToolContext, input string) (string, error) { - roles := toolCtx.AccessChecker.Roles() - requestable := make(map[string]struct{}, 0) - for _, role := range roles { - for _, requestableRole := range role.GetAccessRequestConditions(types.Allow).Roles { - requestable[requestableRole] = struct{}{} - } - } - for _, role := range roles { - for _, requestableRole := range role.GetAccessRequestConditions(types.Deny).Roles { - delete(requestable, requestableRole) - } - } - - resp := strings.Builder{} - for role := range requestable { - resp.Write([]byte(role)) - resp.Write([]byte("\n")) - } - - if resp.Len() == 0 { - return "No requestable roles found", nil - } - - return resp.String(), nil -} - -type AccessRequestListRequestableResourcesTool struct{} - -func (*AccessRequestListRequestableResourcesTool) Name() string { - return "List Requestable Resources" -} - -func (*AccessRequestListRequestableResourcesTool) Description() string { - return `List all resources with IDs that can be requested via access requests. -This includes nodes via SSH access.` -} - -func (a *AccessRequestListRequestableResourcesTool) Run(ctx context.Context, toolCtx *ToolContext, input string) (string, error) { - foundResources := make([]promptResource, 0) - foundResourcesMu := &sync.Mutex{} - g := new(errgroup.Group) - - searchAndAppend := func(resourceType string, convert func(types.Resource) (promptResource, error)) error { - list, err := toolCtx.ListResources(ctx, proto.ListResourcesRequest{ - ResourceType: resourceType, - Limit: maxShownRequestableItems, - UseSearchAsRoles: true, - }) - if err != nil { - return trace.Wrap(err) - } - - for _, resource := range list.Resources { - resource, err := convert(resource) - if err != nil { - return trace.Wrap(err) - } - - foundResourcesMu.Lock() - foundResources = append(foundResources, resource) - foundResourcesMu.Unlock() - } - - return nil - } - searchAndAppendPlain := func(resourceType string) error { - return searchAndAppend(resourceType, func(resource types.Resource) (promptResource, error) { - return promptResource{ - Name: resource.GetName(), - Kind: resource.GetKind(), - SubKind: resource.GetSubKind(), - Labels: resource.GetMetadata().Labels, - }, nil - }) - } - - g.Go(func() error { - return searchAndAppend(types.KindNode, func(resource types.Resource) (promptResource, error) { - return promptResource{ - Name: resource.GetName(), - Kind: resource.GetKind(), - SubKind: resource.GetSubKind(), - Labels: resource.GetMetadata().Labels, - FriendlyName: resource.(types.Server).GetHostname(), - }, nil - }) - }) - g.Go(func() error { return searchAndAppendPlain(types.KindApp) }) - g.Go(func() error { return searchAndAppendPlain(types.KindKubernetesCluster) }) - g.Go(func() error { return searchAndAppendPlain(types.KindDatabase) }) - g.Go(func() error { return searchAndAppendPlain(types.KindWindowsDesktop) }) - - if err := g.Wait(); err != nil { - return "", trace.Wrap(err) - } - sb := strings.Builder{} - total := 0 - for _, resource := range foundResources { - yaml, err := yaml.Marshal(resource) - if err != nil { - return "", trace.Wrap(err) - } - - sb.WriteString(string(yaml)) - sb.WriteString("\n") - - total++ - if total >= maxShownRequestableItems { - break - } - } - - if sb.Len() == 0 { - return "No requestable resources found", nil - } - - return sb.String(), nil -} - -type promptResource struct { - Name string `yaml:"name"` - Kind string `yaml:"kind"` - SubKind string `yaml:"subkind"` - Labels map[string]string `yaml:"labels"` - FriendlyName string `yaml:"friendly_name,omitempty"` -} - -type AccessRequestCreateTool struct{} - -func (*AccessRequestCreateTool) Name() string { - return "Create Access Requests" -} - -func (*AccessRequestCreateTool) Description() string { - return fmt.Sprintf(`Create an access request with a set of roles to, a set of resource UUIDs, a reason, and a set of suggested reviewers. -A valid access request must be either for one or more roles or for one or more resource IDs. -If the user is not specific enough, you may try to determine the correct roles or resource UUIDs by any means you see fit. - -The input must be a JSON object with the following schema: - -%vjson -{ - "roles": []string, \\ The optional set of roles being requested - "resources": []{ - "type": string, \\ The resource type - "id": string, \\ The resource name - "friendlyName": string \\ Optional display-friendly name for the resource - }, \\ The optional set of UUIDs for resources being requested - "reason": string, \\ A reason for the request. This cannot be made up or inferred, it must be explicitly said by the user - "suggested_reviewers": []string \\ An optional list of suggested reviewers; these must be Teleport usernames -} -%v -`, "```", "```") -} - -func (*AccessRequestCreateTool) Run(ctx context.Context, toolCtx *ToolContext, input string) (string, error) { - // This is stubbed because AccessRequestCreateTool is handled specially. - // This is because execution of this tool breaks the loop and returns a suggestion UI prompt. - // It is still handled as a tool because testing has shown that the LLM behaves better when it is treated as a tool. - // - // In addition, treating it as a Tool interface item simplifies the display and prompt assembly logic significantly. - return "", trace.NotImplemented("not implemented") -} - -func (*AccessRequestCreateTool) ParseInput(input string) (*modeloutput.AccessRequest, error) { - output, err := modeloutput.ParseJSONFromModel[modeloutput.AccessRequest](input) - if err != nil { - return nil, err - } - - if output.Reason == "" { - return nil, modeloutput.NewInvalidOutputError( - "access request create: missing reason", - "a reason must be specified for the access request", - ) - } - - if len(output.Roles) == 0 && len(output.Resources) == 0 { - return nil, modeloutput.NewInvalidOutputError( - "access request create: no requested roles or resources", - "an access request must be for one or more roles OR one or more resources", - ) - } - - for i, resource := range output.Resources { - if resource.Type == "" { - return nil, modeloutput.NewInvalidOutputError( - "access request create: missing type at index "+strconv.Itoa(i), - "a type must be provided for each resource", - ) - } - - if resource.Name == "" { - return nil, modeloutput.NewInvalidOutputError( - "access request create: missing name at index "+strconv.Itoa(i), - "a name must be provided for each resource", - ) - } - - if resource.Type == types.KindNode { - if _, err := uuid.Parse(resource.Name); err != nil { - return nil, modeloutput.NewInvalidOutputError( - "access request create: invalid name at index "+strconv.Itoa(i), - "a name must be a valid UUID", - ) - } - } - } - - return &output, nil -} - -type AccessRequestsListTool struct{} - -func (*AccessRequestsListTool) Name() string { - return "List Access Requests" -} - -func (*AccessRequestsListTool) Description() string { - return "List all access requests that the user has access to." -} - -func (*AccessRequestsListTool) Run(ctx context.Context, toolCtx *ToolContext, input string) (string, error) { - requests, err := toolCtx.GetAccessRequests(ctx, types.AccessRequestFilter{ - User: toolCtx.User, - }) - if err != nil { - return "", trace.Wrap(err) - } - - items := make([]accessRequestLLMItem, 0, len(requests)) - for _, request := range requests { - items = append(items, accessRequestLLMItem{ - Roles: request.GetRoles(), - RequestReason: request.GetRequestReason(), - SuggestedReviewers: request.GetSuggestedReviewers(), - State: request.GetState().String(), - ResolveReason: request.GetResolveReason(), - Created: request.GetCreationTime().Format(time.RFC3339), - }) - } - - itemYaml, err := yaml.Marshal(items) - if err != nil { - return "", trace.Wrap(err) - } - - return string(itemYaml), nil -} - -type accessRequestLLMItem struct { - Roles []string `yaml:"roles"` - RequestReason string `yaml:"request_reason"` - SuggestedReviewers []string `yaml:"suggested_reviewers"` - State string `yaml:"state"` - ResolveReason string `yaml:"resolved_reason"` - Created string `yaml:"created"` -} diff --git a/lib/ai/model/tools/auditquery.go b/lib/ai/model/tools/auditquery.go deleted file mode 100644 index 810d12d64c8..00000000000 --- a/lib/ai/model/tools/auditquery.go +++ /dev/null @@ -1,177 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package tools - -import ( - "context" - "fmt" - "strings" - "time" - - "github.com/gravitational/trace" - "github.com/sashabaranov/go-openai" - - "github.com/gravitational/teleport/gen/go/eventschema" - "github.com/gravitational/teleport/lib/ai/model/output" - "github.com/gravitational/teleport/lib/ai/tokens" -) - -const AuditQueryGenerationToolName = "Audit Query Generation" - -type AuditQueryGenerationTool struct { - LLM *openai.Client -} - -func (t *AuditQueryGenerationTool) Name() string { - return AuditQueryGenerationToolName -} - -func (t *AuditQueryGenerationTool) Description() string { - return `Generates a SQL query that can be ran against teleport audit events. -The input must be a single string describing what the query must achieve.` -} - -func (t *AuditQueryGenerationTool) Run(_ context.Context, _ *ToolContext, _ string) (string, error) { - // This is stubbed because AuditQueryGenerationTool is handled specially. - // This is because execution of this tool breaks the loop and returns a command suggestion to the user. - // It is still handled as a tool because testing has shown that the LLM behaves better when it is treated as a tool. - // - // In addition, treating it as a Tool interface item simplifies the display and prompt assembly logic significantly. - return "", trace.NotImplemented("not implemented") -} - -// ChooseEventTable lists all supported events and uses the LLM as a zero shot -// classifier to find which event type can be used to answer the suer query. -func (t *AuditQueryGenerationTool) ChooseEventTable(ctx context.Context, input string, tc *tokens.TokenCount) (string, error) { - tableList, err := eventschema.QueryableEventList() - if err != nil { - return "", trace.Wrap(err) - } - - prompt := []openai.ChatCompletionMessage{ - { - Role: openai.ChatMessageRoleSystem, - Content: `Your job it to find the correct table to run a query on. -You will be given a list of tables, and a request from the user. -You MUST RESPOND ONLY with a single table name. If no table can answer the question, respond 'Cannot answer'.`, - }, - { - Role: openai.ChatMessageRoleUser, - Content: tableList, - }, - { - Role: openai.ChatMessageRoleUser, - Content: fmt.Sprintf("The user request is: %s", input), - }, - } - promptTokens, err := tokens.NewPromptTokenCounter(prompt) - if err != nil { - return "", trace.Wrap(err) - } - tc.AddPromptCounter(promptTokens) - - response, err := t.LLM.CreateChatCompletion( - ctx, - openai.ChatCompletionRequest{ - Model: openai.GPT4, - Messages: prompt, - Temperature: 0, - }, - ) - if err != nil { - return "", trace.Wrap(err) - } - - completion := response.Choices[0].Message.Content - completionTokens, err := tokens.NewSynchronousTokenCounter(completion) - if err != nil { - return "", trace.Wrap(err) - } - tc.AddCompletionCounter(completionTokens) - - eventType := strings.Trim(strings.TrimSpace(strings.ToLower(completion)), "\"'.") - if eventType == "cannot answer" { - return "", trace.NotFound("No relevant event type found. The query cannot be answered by audit logs.") - } - if !eventschema.IsValidEventType(eventType) { - return "", trace.CompareFailed("Model response is not a valid event type: '%s'", eventType) - } - - return eventType, nil - -} - -// GenerateQuery takes an event type, fetches its schema, and calls the LLM to -// generate SQL and answer the user query. -func (t *AuditQueryGenerationTool) GenerateQuery(ctx context.Context, eventType, input string, tc *tokens.TokenCount) (*output.StreamingMessage, error) { - eventSchema, err := eventschema.GetEventSchemaFromType(eventType) - if err != nil { - return nil, trace.Wrap(err) - } - tableSchema, err := eventSchema.TableSchema() - if err != nil { - return nil, trace.Wrap(err) - } - - prompt := []openai.ChatCompletionMessage{ - { - Role: openai.ChatMessageRoleSystem, - Content: fmt.Sprintf(`You are a tool that generates Athena SQL queries to inspect audit events. -You will be given the schema of a table and a user request. -You MUST RESPOND ONLY with an SQL query that answers the user request. -If the request cannot be answered, respond 'none'. -Today's date is DATE('%s')`, time.Now().Format("2006-01-02")), - }, - { - Role: openai.ChatMessageRoleUser, - Content: fmt.Sprintf("The schema of the table `%s` is:\n\n%s", eventschema.SQLViewNameForEvent(eventType), tableSchema), - }, - { - Role: openai.ChatMessageRoleUser, - Content: fmt.Sprintf("The user request is: %s", input), - }, - } - promptTokens, err := tokens.NewPromptTokenCounter(prompt) - if err != nil { - return nil, trace.Wrap(err) - } - tc.AddPromptCounter(promptTokens) - - stream, err := t.LLM.CreateChatCompletionStream( - ctx, - openai.ChatCompletionRequest{ - Model: openai.GPT4, - Messages: prompt, - Temperature: 0, - Stream: true, - }, - ) - if err != nil { - return nil, trace.Wrap(err) - } - - deltas := output.StreamToDeltas(stream) - message, completionTokens, err := output.NewStreamingMessage(deltas, "", "") - if err != nil { - return nil, trace.Wrap(err) - } - tc.AddCompletionCounter(completionTokens) - - return message, nil -} diff --git a/lib/ai/model/tools/commandexec.go b/lib/ai/model/tools/commandexec.go deleted file mode 100644 index 1adc1ffafea..00000000000 --- a/lib/ai/model/tools/commandexec.go +++ /dev/null @@ -1,82 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package tools - -import ( - "context" - "fmt" - - "github.com/gravitational/trace" - - modeloutput "github.com/gravitational/teleport/lib/ai/model/output" -) - -type CommandExecutionTool struct{} - -func (c *CommandExecutionTool) Name() string { - return "Command Execution" -} - -func (c *CommandExecutionTool) Description() string { - return fmt.Sprintf(`Execute a command on a set of remote nodes based on a set of node names or/and a set of labels. -The input must be a JSON object with the following schema: - -%vjson -{ - "command": string, \\ The command to execute - "nodes": []string, \\ Execute a command on all nodes that have the given node names - "labels": []{"key": string, "value": string} \\ Execute a command on all nodes that has at least one of the labels -} -%v -`, "```", "```") -} - -func (c *CommandExecutionTool) Run(_ context.Context, _ *ToolContext, _ string) (string, error) { - // This is stubbed because CommandExecutionTool is handled specially. - // This is because execution of this tool breaks the loop and returns a command suggestion to the user. - // It is still handled as a tool because testing has shown that the LLM behaves better when it is treated as a tool. - // - // In addition, treating it as a Tool interface item simplifies the display and prompt assembly logic significantly. - return "", trace.NotImplemented("not implemented") -} - -// ParseInput is called in a special case if the planned tool is CommandExecutionTool. -// This is because CommandExecutionTool is handled differently from most other tools and forcibly terminates the thought loop. -func (*CommandExecutionTool) ParseInput(input string) (*modeloutput.CompletionCommand, error) { - output, err := modeloutput.ParseJSONFromModel[modeloutput.CompletionCommand](input) - if err != nil { - return nil, err - } - - if output.Command == "" { - return nil, modeloutput.NewInvalidOutputError( - "command execution: missing command", - "command must be non-empty", - ) - } - - if len(output.Nodes) == 0 && len(output.Labels) == 0 { - return nil, modeloutput.NewInvalidOutputError( - "command execution: missing nodes or labels", - "at least one node or label must be specified", - ) - } - - return &output, nil -} diff --git a/lib/ai/model/tools/embedding.go b/lib/ai/model/tools/embedding.go deleted file mode 100644 index f07b8a24490..00000000000 --- a/lib/ai/model/tools/embedding.go +++ /dev/null @@ -1,147 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package tools - -import ( - "context" - "fmt" - "strings" - - "github.com/gravitational/trace" - log "github.com/sirupsen/logrus" - - "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" - "github.com/gravitational/teleport/api/types" - embeddinglib "github.com/gravitational/teleport/lib/ai/embedding" - modeloutput "github.com/gravitational/teleport/lib/ai/model/output" - "github.com/gravitational/teleport/lib/services" -) - -type EmbeddingRetrievalTool struct{} - -type EmbeddingRetrievalToolInput struct { - Question string `json:"question"` -} - -// tryNodeLookupFromProxyCache checks how many nodes the user has access to by -// hitting the proxy cache. If the user has access to less than -// maxEmbeddingsPerLookup, the returned boolean indicates the lookup is -// successful and the result can be used. If the boolean is false, the caller -// must not use the returned result and perform a Node lookup via other means -// (embeddings lookup). -func (e *EmbeddingRetrievalTool) tryNodeLookupFromProxyCache(ctx context.Context, toolCtx *ToolContext) (bool, string, error) { - nodes := toolCtx.NodeWatcher.GetNodes(ctx, func(node services.Node) bool { - err := toolCtx.CheckAccess(node, services.AccessState{MFAVerified: true}) - return err == nil - }) - if len(nodes) == 0 || len(nodes) > maxEmbeddingsPerLookup { - return false, "", nil - } - sb := strings.Builder{} - for _, node := range nodes { - data, err := embeddinglib.SerializeNode(node) - if err != nil { - return false, "", trace.Wrap(err) - } - sb.Write(data) - sb.WriteString("\n") - } - return true, sb.String(), nil -} - -func (e *EmbeddingRetrievalTool) Run(ctx context.Context, toolCtx *ToolContext, input string) (string, error) { - inputCmd, outErr := e.parseInput(input) - if outErr == nil { - // If we failed to parse the input, we can still send the payload for embedding retrieval. - // In most cases, we will still get some sensible results. - // If we parsed the input successfully, we should use the parsed input instead. - input = inputCmd.Question - } - log.Tracef("embedding retrieval input: %v", input) - - // Threshold to avoid looping over all nodes on large clusters - if toolCtx.NodeWatcher != nil && toolCtx.NodeWatcher.NodeCount() < proxyLookupClusterMaxSize { - ok, result, err := e.tryNodeLookupFromProxyCache(ctx, toolCtx) - if err != nil { - return "", trace.Wrap(err) - } - if ok { - return result, nil - } - } - - resp, err := toolCtx.GetAssistantEmbeddings(ctx, &assist.GetAssistantEmbeddingsRequest{ - Username: toolCtx.User, - Kind: types.KindNode, // currently only node embeddings are supported - Limit: maxEmbeddingsPerLookup, - Query: input, - }) - if err != nil { - return "", trace.Wrap(err) - } - - sb := strings.Builder{} - for _, embedding := range resp.Embeddings { - sb.WriteString(embedding.Content) - sb.WriteString("\n") - } - - log.Tracef("embedding retrieval: %v", sb.String()) - - if sb.Len() == 0 { - // Either no nodes are connected, embedding process hasn't started yet, or - // the user doesn't have access to any resources. - return "Didn't find any nodes matching the query", nil - } - - return sb.String(), nil -} - -func (e *EmbeddingRetrievalTool) Name() string { - return "Nodes names and labels retrieval" -} - -func (e *EmbeddingRetrievalTool) Description() string { - return fmt.Sprintf(`Ask about existing remote nodes that user has access to fetch node names or/and set of labels. -Always use this capability before returning generating any command. Do not assume that the user has access to any nodes. Returning a command without checking for access will result in an error. -Always prefer to use labler rather than node names. -The input must be a JSON object with the following schema: -%vjson -{ - "question": string \\ Question about the available remote nodes -} -%v -`, "```", "```") -} - -func (*EmbeddingRetrievalTool) parseInput(input string) (*EmbeddingRetrievalToolInput, error) { - output, err := modeloutput.ParseJSONFromModel[EmbeddingRetrievalToolInput](input) - if err != nil { - return nil, err - } - - if len(output.Question) == 0 { - return nil, modeloutput.NewInvalidOutputError( - "embedding retrieval: missing question", - "question must be non-empty", - ) - } - - return &output, nil -} diff --git a/lib/ai/model/tools/embedding_test.go b/lib/ai/model/tools/embedding_test.go deleted file mode 100644 index 5214e61851e..00000000000 --- a/lib/ai/model/tools/embedding_test.go +++ /dev/null @@ -1,140 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package tools - -import ( - "context" - "fmt" - "testing" - - "github.com/gravitational/trace" - "github.com/stretchr/testify/require" - - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/lib/services" -) - -const testUser = "username" - -// mockAccessChecker implements the services.AccessChecker and always validate -// or reject access based on its allowAccess field. -type mockAccessChecker struct { - allowAccess bool - services.AccessChecker -} - -func (ac *mockAccessChecker) CheckAccess(_ services.AccessCheckable, _ services.AccessState, _ ...services.RoleMatcher) error { - if ac.allowAccess { - return nil - } - return trace.AccessDenied("user does not have access") -} - -// mockNodeGetter returns a static list of nodes -type mockNodeGetter struct { - nodes []types.Server -} - -func (ng *mockNodeGetter) NodeCount() int { - return len(ng.nodes) -} - -func (ng *mockNodeGetter) GetNodes(_ context.Context, fn func(n services.Node) bool) []types.Server { - var result []types.Server - for _, node := range ng.nodes { - if fn(node) { - result = append(result, node) - } - } - return result -} - -func Test_embeddingRetrievalTool_tryNodeLookupFromProxyCache(t *testing.T) { - ctx := context.Background() - tests := []struct { - name string - nodeCount int - hasAccess bool - assertLookupSuccessful require.BoolAssertionFunc - expectedOutput string - }{ - { - name: "No nodes", - nodeCount: 0, - hasAccess: true, - assertLookupSuccessful: require.False, - }, - { - name: "Few nodes", - nodeCount: 2, - hasAccess: true, - assertLookupSuccessful: require.True, - expectedOutput: `name: node-0 -kind: node -subkind: teleport -labels: - foo: bar - -name: node-1 -kind: node -subkind: teleport -labels: - foo: bar - -`, - }, - { - name: "Few nodes without access", - nodeCount: 2, - hasAccess: false, - assertLookupSuccessful: require.False, - }, - { - name: "Too many nodes", - nodeCount: maxEmbeddingsPerLookup + 1, - hasAccess: true, - assertLookupSuccessful: require.False, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - // Test setup - var err error - nodes := make([]types.Server, tt.nodeCount) - for i := 0; i < tt.nodeCount; i++ { - nodeName := fmt.Sprintf("node-%d", i) - nodes[i], err = types.NewServerWithLabels(nodeName, types.KindNode, types.ServerSpecV2{Hostname: nodeName}, map[string]string{"foo": "bar"}) - require.NoError(t, err) - } - - toolCtx := &ToolContext{ - User: testUser, - AccessChecker: &mockAccessChecker{allowAccess: tt.hasAccess}, - NodeWatcher: &mockNodeGetter{nodes: nodes}, - } - - // Doing the real test - tool := EmbeddingRetrievalTool{} - ok, output, err := tool.tryNodeLookupFromProxyCache(ctx, toolCtx) - require.NoError(t, err) - tt.assertLookupSuccessful(t, ok) - require.Equal(t, tt.expectedOutput, output) - }) - } -} diff --git a/lib/ai/model/tools/generationtool.go b/lib/ai/model/tools/generationtool.go deleted file mode 100644 index 19c6ff2a52c..00000000000 --- a/lib/ai/model/tools/generationtool.go +++ /dev/null @@ -1,84 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package tools - -import ( - "context" - "fmt" - - "github.com/gravitational/trace" - - modeloutput "github.com/gravitational/teleport/lib/ai/model/output" -) - -type CommandGenerationTool struct{} - -type CommandGenerationToolInput struct { - // Command is a unix command to execute. - Command string `json:"command"` -} - -func (c *CommandGenerationTool) Name() string { - return "Command Generation" -} - -func (c *CommandGenerationTool) Description() string { - // acknowledgement field is used to convince the LLM to return the JSON. - // Base on my testing LLM ignores the JSON when the schema has only one field. - // Adding additional "pseudo-fields" to the schema makes the LLM return the JSON. - return fmt.Sprintf(`Generate a Bash command. -The input must be a JSON object with the following schema: -%vjson -{ - "command": string, \\ The generated command - "acknowledgement": boolean \\ Set to true to ackowledge that you understand the formatting -} -%v -`, "```", "```") -} - -func (c *CommandGenerationTool) Run(_ context.Context, toolCtx *ToolContext, _ string) (string, error) { - // This is stubbed because CommandGenerationTool is handled specially. - // This is because execution of this tool breaks the loop and returns a command suggestion to the user. - // It is still handled as a tool because testing has shown that the LLM behaves better when it is treated as a tool. - // - // In addition, treating it as a Tool interface item simplifies the display and prompt assembly logic significantly. - return "", trace.NotImplemented("not implemented") -} - -// ParseInput is called in a special case if the planned tool is CommandExecutionTool. -// This is because CommandExecutionTool is handled differently from most other tools and forcibly terminates the thought loop. -func (*CommandGenerationTool) ParseInput(input string) (*CommandGenerationToolInput, error) { - output, err := modeloutput.ParseJSONFromModel[CommandGenerationToolInput](input) - if err != nil { - return nil, err - } - - if output.Command == "" { - return nil, modeloutput.NewInvalidOutputError( - "command generation: missing command", - "command must be non-empty", - ) - } - - // Ignore the acknowledgement field. - // We do not care about the value. Having the command it enough. - - return &output, nil -} diff --git a/lib/ai/model/tools/tool.go b/lib/ai/model/tools/tool.go deleted file mode 100644 index 34ac5527ff2..00000000000 --- a/lib/ai/model/tools/tool.go +++ /dev/null @@ -1,77 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package tools - -import ( - "context" - - "github.com/gravitational/teleport/api/client/proto" - "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/lib/services" -) - -const ( - // proxyLookupClusterMaxSize is max the number of nodes in the cluster to attempt an opportunistic node lookup - // in the proxy cache. We always do embedding lookups if the cluster is larger than this number. - proxyLookupClusterMaxSize = 100 - maxEmbeddingsPerLookup = 10 - - // TODO(joel): remove/change when migrating to embeddings - maxShownRequestableItems = 50 -) - -// *ToolContext contains various "data" which is commonly needed by various tools. -type ToolContext struct { - assist.AssistEmbeddingServiceClient - AccessRequestClient - AccessPoint - services.AccessChecker - NodeWatcher NodeWatcher - User string - ClusterName string -} - -// NodeWatcher abstracts away services.NodeWatcher for testing purposes. -type NodeWatcher interface { - // GetNodes returns a list of nodes that match the given filter. - GetNodes(ctx context.Context, fn func(n services.Node) bool) []types.Server - - // NodeCount returns the number of nodes in the cluster. - NodeCount() int -} - -// AccessPoint allows reading resources from proxy cache. -type AccessPoint interface { - ListResources(ctx context.Context, req proto.ListResourcesRequest) (*types.ListResourcesResponse, error) -} - -// AccessRequestClient abstracts away the access request client for testing purposes. -type AccessRequestClient interface { - CreateAccessRequestV2(ctx context.Context, req types.AccessRequest) (types.AccessRequest, error) - GetAccessRequests(ctx context.Context, filter types.AccessRequestFilter) ([]types.AccessRequest, error) -} - -// Tool is an interface that allows the agent to interact with the outside world. -// It is used to implement things such as vector document retrieval and command execution. -type Tool interface { - Name() string - Description() string - Run(ctx context.Context, toolCtx *ToolContext, input string) (string, error) -} diff --git a/lib/ai/simpleretriever.go b/lib/ai/simpleretriever.go deleted file mode 100644 index 8e909638ecf..00000000000 --- a/lib/ai/simpleretriever.go +++ /dev/null @@ -1,146 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package ai - -import ( - "sort" - "sync" - - "github.com/gravitational/trace" - - "github.com/gravitational/teleport/lib/ai/embedding" -) - -// Document is a embedding enriched with similarity score -type Document struct { - *embedding.Embedding - SimilarityScore float64 -} - -func calculateSimilarity(v1, v2 []float64) (float64, error) { - if len(v1) != len(v2) { - return 0, trace.BadParameter("vectors must be the same length") - } - - var result float64 - for i, val := range v1 { - result += val * v2[i] - } - - return result, nil -} - -// SimpleRetriever is a simple implementation of embeddings retriever. -// It stores all the embeddings in memory and retrieves the k nearest neighbors -// by iterating over all the embeddings. Do not use for large datasets. -type SimpleRetriever struct { - embeddings map[string]*embedding.Embedding - maxSize int - mtx sync.Mutex -} - -func NewSimpleRetriever() *SimpleRetriever { - return &SimpleRetriever{ - embeddings: make(map[string]*embedding.Embedding), - maxSize: 1_000, // keep the number low to avoid OOM - } -} - -// Insert adds the embedding to the retriever. If the retriever is full, the -// embedding is not added and false is returned. -func (r *SimpleRetriever) Insert(id string, embedding *embedding.Embedding) bool { - r.mtx.Lock() - defer r.mtx.Unlock() - if len(r.embeddings) >= r.maxSize { - return false - } - r.embeddings[id] = embedding - return true -} - -// Remove removes the embedding from the retriever by ID. -func (r *SimpleRetriever) Remove(id string) { - r.mtx.Lock() - defer r.mtx.Unlock() - delete(r.embeddings, id) -} - -// Swap replaces the embeddings in the retriever with the embeddings from the -// provided retriever. -// The mutex is acquired for the receiver, but not for the provided retriever. -func (r *SimpleRetriever) Swap(s *SimpleRetriever) { - r.mtx.Lock() - defer r.mtx.Unlock() - r.embeddings = s.embeddings - r.maxSize = s.maxSize -} - -// FilterFn is a function that filters out embeddings. -// If the function returns false, the embedding is filtered out. -type FilterFn func(id string, embedding *embedding.Embedding) bool - -// GetRelevant returns the k nearest neighbors to the query embedding. -// If a filter is provided, only the embeddings that pass the filter are considered. -func (r *SimpleRetriever) GetRelevant(query *embedding.Embedding, k int, filter FilterFn) []*Document { - // Replace with priority queue if k is large. - results := make([]*Document, 0, k) - - r.mtx.Lock() - defer r.mtx.Unlock() - - // Find the k nearest neighbors - for id, embedding := range r.embeddings { - // Skip if the document is filtered out - if filter != nil && !filter(id, embedding) { - continue - } - - // Calculate the similarity score - similarity, _ := calculateSimilarity(query.Vector, embedding.Vector) - // If the results slice smaller than the k, add the element to the results - if len(results) < k { - results = append(results, &Document{ - Embedding: embedding, - SimilarityScore: similarity, - }) - - // Sort to preserve the invariant - the result slice is sorted by - // similarity score - sort.Slice(results, func(i, j int) bool { - return results[i].SimilarityScore > results[j].SimilarityScore - }) - } else if similarity > results[len(results)-1].SimilarityScore { - // If the element is more relevant than the least similar element, - // add it to the result slice - results[len(results)-1] = &Document{ - Embedding: embedding, - SimilarityScore: similarity, - } - - // Sort to preserve the invariant - the result slice is sorted by - // similarity score - sort.Slice(results, func(i, j int) bool { - return results[i].SimilarityScore > results[j].SimilarityScore - }) - } - } - - // Return the results sorted by similarity score. - return results -} diff --git a/lib/ai/simpleretriever_test.go b/lib/ai/simpleretriever_test.go deleted file mode 100644 index c05cee1a8ff..00000000000 --- a/lib/ai/simpleretriever_test.go +++ /dev/null @@ -1,103 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package ai - -import ( - "fmt" - "math" - "math/rand" - "strconv" - "testing" - - "github.com/stretchr/testify/require" - - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/lib/ai/embedding" -) - -// Function to calculate L2 norm -func L2norm(v []float64) float64 { - sum := 0.0 - for _, value := range v { - sum += value * value - } - return math.Sqrt(sum) -} - -// Function to normalize vector using L2 norm -func normalize(v embedding.Vector64) embedding.Vector64 { - norm := L2norm(v) - result := make(embedding.Vector64, len(v)) - for i, value := range v { - result[i] = value / norm - } - return result -} - -func TestSimpleRetriever_GetRelevant(t *testing.T) { - t.Parallel() - - // Generate random vector. The seed is fixed, so the results are deterministic. - randGen := rand.New(rand.NewSource(42)) - - generateVector := func() embedding.Vector64 { - const testVectorDimension = 100 - // generate random vector - // reduce the dimensionality to 100 - vec := make(embedding.Vector64, testVectorDimension) - for i := 0; i < testVectorDimension; i++ { - vec[i] = randGen.Float64() - } - // normalize vector, so the similarity between two vectors is the dot product - // between [0, 1] - return normalize(vec) - } - - const testEmbeddingsSize = 100 - points := make([]*embedding.Embedding, testEmbeddingsSize) - for i := 0; i < testEmbeddingsSize; i++ { - points[i] = embedding.NewEmbedding(types.KindNode, strconv.Itoa(i), generateVector(), [32]byte{}) - } - - // Create a query. - query := embedding.NewEmbedding(types.KindNode, "1", generateVector(), [32]byte{}) - - retriever := NewSimpleRetriever() - - for _, point := range points { - retriever.Insert(point.GetName(), point) - } - - // Get the top 10 most similar documents. - docs := retriever.GetRelevant(query, 10, func(id string, embedding *embedding.Embedding) bool { - return true - }) - require.Len(t, docs, 10) - - expectedResults := []int{57, 92, 95, 49, 33, 56, 30, 99, 90, 47} - expectedSimilarities := []float64{0.80405, 0.79051, 0.78161, 0.78159, - 0.77655, 0.77374, 0.77306, 0.76688, 0.76634, 0.76458} - - for i, result := range docs { - require.Equal(t, - fmt.Sprintf("%s/%s", types.KindNode, strconv.Itoa(expectedResults[i])), - result.GetName(), "expected order is wrong") - require.InDelta(t, expectedSimilarities[i], result.SimilarityScore, 10e-6, "similarity score is wrong") - } -} diff --git a/lib/ai/testutils/http.go b/lib/ai/testutils/http.go deleted file mode 100644 index 6cfabeae20a..00000000000 --- a/lib/ai/testutils/http.go +++ /dev/null @@ -1,131 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package testutils - -import ( - "encoding/json" - "net/http" - "strconv" - "testing" - "time" - - "github.com/sashabaranov/go-openai" - "github.com/stretchr/testify/assert" -) - -// GetTestHandlerFn returns a handler function that can be used to OpenAI API used by -// the chat API. It takes a list of responses that will be returned in order. -func GetTestHandlerFn(t *testing.T, responses []string) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPost || !(r.URL.Path == "/chat/completions") { - http.Error(w, "Unexpected request", http.StatusBadRequest) - return - } - - switch r.Header.Get("Accept") { - case "application/json; charset=utf-8", "application/json": - responses = messageResponse(w, r, t, responses) - case "text/event-stream": - responses = streamResponse(w, t, responses) - default: - http.Error(w, "Unexpected request", http.StatusBadRequest) - } - } -} - -func streamResponse(w http.ResponseWriter, t *testing.T, responses []string) []string { - w.Header().Set("Content-Type", "text/event-stream") - - if !assert.NotEmpty(t, responses, "Unexpected request") { - http.Error(w, "Unexpected request", http.StatusBadRequest) - return responses - } - - resp := &openai.ChatCompletionStreamResponse{ - ID: strconv.Itoa(int(time.Now().Unix())), - Object: "completion", - Created: time.Now().Unix(), - Model: openai.GPT4, - Choices: []openai.ChatCompletionStreamChoice{ - { - Index: 0, - Delta: openai.ChatCompletionStreamChoiceDelta{ - Content: responses[0], - Role: openai.ChatMessageRoleAssistant, - }, - FinishReason: "", - }, - }, - } - - respBytes, err := json.Marshal(resp) - assert.NoError(t, err, "Marshal error") - - _, err = w.Write([]byte("data: ")) - assert.NoError(t, err, "Write error") - _, err = w.Write(respBytes) - assert.NoError(t, err, "Write error") - _, err = w.Write([]byte("\n\nevent: done\ndata: [DONE]\n\n")) - assert.NoError(t, err, "Write error") - - return responses[1:] -} - -func messageResponse(w http.ResponseWriter, r *http.Request, t *testing.T, responses []string) []string { - w.Header().Set("Content-Type", "application/json") - - req := &openai.ChatCompletionRequest{} - err := json.NewDecoder(r.Body).Decode(req) - if err != nil { - http.Error(w, err.Error(), http.StatusBadRequest) - } - - // Use assert as require doesn't work when called from a goroutine - if !assert.NotEmpty(t, responses, "Unexpected request") { - http.Error(w, "Unexpected request", http.StatusBadRequest) - return responses - } - - dataBytes := responses[0] - - resp := openai.ChatCompletionResponse{ - ID: strconv.Itoa(int(time.Now().Unix())), - Object: "test-object", - Created: time.Now().Unix(), - Model: req.Model, - Choices: []openai.ChatCompletionChoice{ - { - Message: openai.ChatCompletionMessage{ - Role: openai.ChatMessageRoleAssistant, - Content: dataBytes, - Name: "", - }, - }, - }, - Usage: openai.Usage{}, - } - - respBytes, err := json.Marshal(resp) - assert.NoError(t, err, "Marshal error") - - _, err = w.Write(respBytes) - assert.NoError(t, err, "Write error") - - return responses[1:] -} diff --git a/lib/ai/tokens/contants.go b/lib/ai/tokens/contants.go deleted file mode 100644 index 80ec66d7e15..00000000000 --- a/lib/ai/tokens/contants.go +++ /dev/null @@ -1,31 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package tokens - -// Ref: https://github.com/openai/openai-cookbook/blob/594fc6c952425810e9ea5bd1a275c8ca5f32e8f9/examples/How_to_count_tokens_with_tiktoken.ipynb -const ( - // perMessage is the token "overhead" for each message - perMessage = 3 - - // perRequest is the number of tokens used up for each completion request - perRequest = 3 - - // perRole is the number of tokens used to encode a message's role - perRole = 1 -) diff --git a/lib/ai/tokens/tokencount.go b/lib/ai/tokens/tokencount.go deleted file mode 100644 index 8550d2aae36..00000000000 --- a/lib/ai/tokens/tokencount.go +++ /dev/null @@ -1,194 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package tokens - -import ( - "sync" - - "github.com/gravitational/trace" - "github.com/sashabaranov/go-openai" -) - -// TokenCount holds TokenCounters for both Prompt and Completion tokens. -// As the agent performs multiple calls to the model, each call creates its own -// prompt and completion TokenCounter. -// -// Prompt TokenCounters can be created before doing the call as we know the -// full prompt and can tokenize it. This is the PromptTokenCounter purpose. -// -// Completion TokenCounters can be created after receiving the model response. -// Depending on the response type, we might have the full result already or get -// a stream that will provide the completion result in the future. For the latter, -// the token count will be evaluated lazily and asynchronously. -// StaticTokenCounter count tokens synchronously, while -// AsynchronousTokenCounter supports the streaming use-cases. -type TokenCount struct { - Prompt TokenCounters - Completion TokenCounters -} - -// AddPromptCounter adds a TokenCounter to the Prompt list. -func (tc *TokenCount) AddPromptCounter(prompt TokenCounter) { - if prompt != nil { - tc.Prompt = append(tc.Prompt, prompt) - } -} - -// AddCompletionCounter adds a TokenCounter to the Completion list. -func (tc *TokenCount) AddCompletionCounter(completion TokenCounter) { - if completion != nil { - tc.Completion = append(tc.Completion, completion) - } -} - -// CountAll iterates over all counters and returns how many prompt and -// completion tokens were used. As completion token counting can require waiting -// for a response to be streamed, the caller should pass a context and use it to -// implement some kind of deadline to avoid hanging infinitely if something goes -// wrong (e.g. use `context.WithTimeout()`). -func (tc *TokenCount) CountAll() (int, int) { - return tc.Prompt.CountAll(), tc.Completion.CountAll() -} - -// NewTokenCount initializes a new TokenCount struct. -func NewTokenCount() *TokenCount { - return &TokenCount{ - Prompt: TokenCounters{}, - Completion: TokenCounters{}, - } -} - -// TokenCounter is an interface for all token counters, regardless of the kind -// of token they count (prompt/completion) or the tokenizer used. -// TokenCount must be idempotent. -type TokenCounter interface { - TokenCount() int -} - -// TokenCounters is a list of TokenCounter and offers function to iterate over -// all counters and compute the total. -type TokenCounters []TokenCounter - -// CountAll iterates over a list of TokenCounter and returns the sum of the -// results of all counters. As the counting process might be blocking/take some -// time, the caller should set a Deadline on the context. -func (tc TokenCounters) CountAll() int { - var total int - for _, counter := range tc { - total += counter.TokenCount() - } - return total -} - -// StaticTokenCounter is a token counter whose count has already been evaluated. -// This can be used to count prompt tokens (we already know the exact count), -// or to count how many tokens were used by an already finished completion -// request. -type StaticTokenCounter int - -// TokenCount implements the TokenCounter interface. -func (tc *StaticTokenCounter) TokenCount() int { - return int(*tc) -} - -// NewPromptTokenCounter takes a list of openai.ChatCompletionMessage and -// computes how many tokens are used by sending those messages to the model. -func NewPromptTokenCounter(prompt []openai.ChatCompletionMessage) (*StaticTokenCounter, error) { - var promptCount int - for _, message := range prompt { - promptTokens := countTokens(message.Content) - - promptCount = promptCount + perMessage + perRole + promptTokens - } - tc := StaticTokenCounter(promptCount) - - return &tc, nil -} - -// NewSynchronousTokenCounter takes the completion request output and -// computes how many tokens were used by the model to generate this result. -func NewSynchronousTokenCounter(completion string) (*StaticTokenCounter, error) { - completionTokens := countTokens(completion) - completionCount := perRequest + completionTokens - - tc := StaticTokenCounter(completionCount) - return &tc, nil -} - -// AsynchronousTokenCounter counts completion tokens that are used by a -// streamed completion request. When creating a AsynchronousTokenCounter, -// the streaming might not be finished, and we can't evaluate how many tokens -// will be used. In this case, the streaming routine must add streamed -// completion result with the Add() method and call Finish() once the -// completion is finished. TokenCount() will hang until either Finish() is -// called or the context is Done. -type AsynchronousTokenCounter struct { - count int - - // mutex protects all fields of the AsynchronousTokenCounter, it must be - // acquired before any read or write operation. - mutex sync.Mutex - // finished tells if the count is finished or not. - // TokenCount() finishes the count. Once the count is finished, Add() will - // throw errors. - finished bool -} - -// TokenCount implements the TokenCounter interface. -// It returns how many tokens have been counted. It also marks the counter as -// finished. Once a counter is finished, tokens cannot be added anymore. -func (tc *AsynchronousTokenCounter) TokenCount() int { - // If the count is already finished, we return the values - tc.mutex.Lock() - defer tc.mutex.Unlock() - tc.finished = true - return tc.count + perRequest -} - -// Add a streamed token to the count. -func (tc *AsynchronousTokenCounter) Add() error { - tc.mutex.Lock() - defer tc.mutex.Unlock() - - if tc.finished { - return trace.Errorf("Count is already finished, cannot add more content") - } - tc.count += 1 - return nil -} - -// NewAsynchronousTokenCounter takes the partial completion request output -// and creates a token counter that can be already returned even if not all -// the content has been streamed yet. Streamed content can be added a posteriori -// with Add(). Once all the content is streamed, Finish() must be called. -func NewAsynchronousTokenCounter(completionStart string) (*AsynchronousTokenCounter, error) { - completionTokens := countTokens(completionStart) - - return &AsynchronousTokenCounter{ - count: completionTokens, - mutex: sync.Mutex{}, - finished: false, - }, nil -} - -// countTokens returns an estimated number of tokens in the text. -func countTokens(text string) int { - // Rough estimations that each token is around 4 characters. - return len(text) / 4 -} diff --git a/lib/ai/tokens/tokencount_test.go b/lib/ai/tokens/tokencount_test.go deleted file mode 100644 index c584ef2ba5c..00000000000 --- a/lib/ai/tokens/tokencount_test.go +++ /dev/null @@ -1,96 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package tokens - -import ( - "testing" - - "github.com/stretchr/testify/require" -) - -const ( - testCompletionStart = "This is the beginning of the response." - testCompletionEnd = "And this is the end." -) - -// This test checks that Add() properly appends content in the completion -// response. -func TestAsynchronousTokenCounter_TokenCount(t *testing.T) { - t.Parallel() - tests := []struct { - name string - completionStart string - completionEnd string - expectedTokens int - }{ - { - name: "empty count", - expectedTokens: 3, - }, - { - name: "only completion start", - completionStart: testCompletionStart, - expectedTokens: 12, - }, - { - name: "only completion add", - completionEnd: testCompletionEnd, - expectedTokens: 8, - }, - { - name: "completion start and end", - completionStart: testCompletionStart, - completionEnd: testCompletionEnd, - expectedTokens: 17, - }, - } - for _, tt := range tests { - tt := tt - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - // Test setup - tc, err := NewAsynchronousTokenCounter(tt.completionStart) - require.NoError(t, err) - tokens := countTokens(tt.completionEnd) - - for range tokens { - require.NoError(t, tc.Add()) - } - - // Doing the real test: asserting the count is right - count := tc.TokenCount() - require.Equal(t, tt.expectedTokens, count) - }) - } -} - -func TestAsynchronousTokenCounter_Finished(t *testing.T) { - tc, err := NewAsynchronousTokenCounter(testCompletionStart) - require.NoError(t, err) - - // We can Add() if the counter has not been read yet - require.NoError(t, tc.Add()) - - // We read from the counter - count := tc.TokenCount() - require.Equal(t, 13, count) - - // Adding new tokens should be impossible - require.Error(t, tc.Add()) -} diff --git a/lib/assist/assist.go b/lib/assist/assist.go deleted file mode 100644 index 7beefc48b15..00000000000 --- a/lib/assist/assist.go +++ /dev/null @@ -1,623 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package assist - -import ( - "context" - "encoding/json" - "strings" - "time" - - "github.com/gravitational/trace" - "github.com/gravitational/trace/trail" - "github.com/jonboulle/clockwork" - "github.com/sashabaranov/go-openai" - log "github.com/sirupsen/logrus" - "google.golang.org/protobuf/types/known/timestamppb" - - "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" - pluginsv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/plugins/v1" - "github.com/gravitational/teleport/lib/ai" - "github.com/gravitational/teleport/lib/ai/model" - "github.com/gravitational/teleport/lib/ai/model/output" - "github.com/gravitational/teleport/lib/ai/model/tools" - "github.com/gravitational/teleport/lib/ai/tokens" -) - -// MessageType is a type of the Assist message. -type MessageType string - -const ( - // MessageKindCommand is the type of Assist message that contains the command to execute. - MessageKindCommand MessageType = "COMMAND" - // MessageKindCommandResult is the type of Assist message that contains the command execution result. - MessageKindCommandResult MessageType = "COMMAND_RESULT" - // MessageKindAccessRequest is the type of Assist message that contains the access request. - // Sent by the backend when it wants the frontend to display a prompt to the user. - MessageKindAccessRequest MessageType = "ACCESS_REQUEST" - // MessageKindAccessRequestCreated is a marker message to indicate that an access request was created. - // Sent by the frontend to the backend to indicate that it was created to future loads of the conversation. - MessageKindAccessRequestCreated MessageType = "ACCESS_REQUEST_CREATED" - // MessageKindCommandResultSummary is the type of message that is optionally - // emitted after a command and contains a summary of the command output. - // This message is both sent after the command execution to the web UI, - // and persisted in the conversation history. - MessageKindCommandResultSummary MessageType = "COMMAND_RESULT_SUMMARY" - // MessageKindUserMessage is the type of Assist message that contains the user message. - MessageKindUserMessage MessageType = "CHAT_MESSAGE_USER" - // MessageKindAssistantMessage is the type of Assist message that contains the assistant message. - MessageKindAssistantMessage MessageType = "CHAT_MESSAGE_ASSISTANT" - // MessageKindAssistantPartialMessage is the type of Assist message that contains the assistant partial message. - MessageKindAssistantPartialMessage MessageType = "CHAT_PARTIAL_MESSAGE_ASSISTANT" - // MessageKindAssistantPartialFinalize is the type of Assist message that ends the partial message stream. - MessageKindAssistantPartialFinalize MessageType = "CHAT_PARTIAL_MESSAGE_ASSISTANT_FINALIZE" - // MessageKindSystemMessage is the type of Assist message that contains the system message. - MessageKindSystemMessage MessageType = "CHAT_MESSAGE_SYSTEM" - // MessageKindError is the type of Assist message that is presented to user as information, but not stored persistently in the conversation. This can include backend error messages and the like. - MessageKindError MessageType = "CHAT_MESSAGE_ERROR" - // MessageKindProgressUpdate is the type of Assist message that contains a progress update. - // A progress update starts a new "stage" and ends a previous stage if there was one. - MessageKindProgressUpdate MessageType = "CHAT_MESSAGE_PROGRESS_UPDATE" -) - -// PluginGetter is the minimal interface used by the chat to interact with the plugin service in the backend. -type PluginGetter interface { - PluginsClient() pluginsv1.PluginServiceClient -} - -// MessageService is the minimal interface used by the chat to interact with the Assist message service in the backend. -type MessageService interface { - // GetAssistantMessages returns all messages with given conversation ID. - GetAssistantMessages(ctx context.Context, req *assist.GetAssistantMessagesRequest) (*assist.GetAssistantMessagesResponse, error) - - // CreateAssistantMessage adds the message to the backend. - CreateAssistantMessage(ctx context.Context, msg *assist.CreateAssistantMessageRequest) error -} - -// Assist is the Teleport Assist client. -type Assist struct { - client *ai.Client - // clock is a clock used to generate timestamps. - clock clockwork.Clock -} - -// NewClient creates a new Assist client. -func NewClient(ctx context.Context, proxyClient PluginGetter, - proxySettings any, openaiCfg *openai.ClientConfig) (*Assist, error) { - - client, err := getAssistantClient(ctx, proxyClient, proxySettings, openaiCfg) - if err != nil { - return nil, trace.Wrap(err) - } - - return &Assist{ - client: client, - clock: clockwork.NewRealClock(), - }, nil -} - -// Chat is a Teleport Assist chat. -type Chat struct { - assist *Assist - chat *ai.Chat - // assistService is the auth server client. - assistService MessageService - // ConversationID is the ID of the conversation. - ConversationID string - // Username is the username of the user who started the chat. - Username string - // potentiallyStaleHistory indicates messages might have been inserted into - // the chat history and the messages should be re-fetched before attempting - // the next completion. - potentiallyStaleHistory bool -} - -// NewChat creates a new Assist chat. -func (a *Assist) NewChat(ctx context.Context, assistService MessageService, toolContext *tools.ToolContext, - conversationID string, -) (*Chat, error) { - aichat := a.client.NewChat(toolContext) - - chat := &Chat{ - assist: a, - assistService: assistService, - chat: aichat, - ConversationID: conversationID, - Username: toolContext.User, - potentiallyStaleHistory: false, - } - - if err := chat.loadMessages(ctx); err != nil { - return nil, trace.Wrap(err) - } - - return chat, nil -} - -// LightweightChat is a Teleport Assist chat that doesn't store the history -// of the conversation. -type LightweightChat struct { - assist *Assist - chat *ai.Chat -} - -// NewLightweightChat creates a new Assist chat what doesn't store the history -// of the conversation. -func (a *Assist) NewLightweightChat(username string) (*LightweightChat, error) { - aichat := a.client.NewCommand(username) - return &LightweightChat{ - assist: a, - chat: aichat, - }, nil -} - -func (a *Assist) NewSSHCommand(username string) (*ai.Chat, error) { - return a.client.NewCommand(username), nil -} - -// GenerateSummary generates a summary for the given message. -func (a *Assist) GenerateSummary(ctx context.Context, message string) (string, error) { - return a.client.Summary(ctx, message) -} - -// RunTool runs a model tool without an ai.Chat. -func (a *Assist) RunTool(ctx context.Context, onMessage onMessageFunc, toolName, userInput string, toolContext *tools.ToolContext, -) (*tokens.TokenCount, error) { - message, tc, err := a.client.RunTool(ctx, toolContext, toolName, userInput) - if err != nil { - return nil, trace.Wrap(err) - } - - switch message := message.(type) { - case *output.Message: - if err := onMessage(MessageKindAssistantMessage, []byte(message.Content), a.clock.Now().UTC()); err != nil { - return nil, trace.Wrap(err) - } - case *output.GeneratedCommand: - if err := onMessage(MessageKindCommand, []byte(message.Command), a.clock.Now().UTC()); err != nil { - return nil, trace.Wrap(err) - } - case *output.StreamingMessage: - if err := func() error { - var text strings.Builder - defer onMessage(MessageKindAssistantPartialFinalize, nil, a.clock.Now().UTC()) - for part := range message.Parts { - text.WriteString(part) - - if err := onMessage(MessageKindAssistantPartialMessage, []byte(part), a.clock.Now().UTC()); err != nil { - return trace.Wrap(err) - } - } - return nil - }(); err != nil { - return nil, trace.Wrap(err) - } - default: - return nil, trace.Errorf("Unexpected message type: %T", message) - } - - return tc, nil -} - -// GenerateCommandSummary summarizes the output of a command executed on one or -// many nodes. The conversation history is also sent into the prompt in order -// to gather context and know what information is relevant in the command output. -func (a *Assist) GenerateCommandSummary(ctx context.Context, messages []*assist.AssistantMessage, output map[string][]byte) (string, *tokens.TokenCount, error) { - // Create system prompt - modelMessages := []openai.ChatCompletionMessage{ - {Role: openai.ChatMessageRoleSystem, Content: model.PromptSummarizeCommand}, - } - - // Load context back into prompt - for _, message := range messages { - role := kindToRole(MessageType(message.Type)) - if role != "" && role != openai.ChatMessageRoleSystem { - payload, err := formatMessagePayload(message) - if err != nil { - return "", nil, trace.Wrap(err) - } - modelMessages = append(modelMessages, openai.ChatCompletionMessage{Role: role, Content: payload}) - } - } - return a.client.CommandSummary(ctx, modelMessages, output) -} - -// reloadMessages clears the chat history and reloads the messages from the database. -func (c *Chat) reloadMessages(ctx context.Context) error { - c.chat.Clear() - return c.loadMessages(ctx) -} - -// ClassifyMessage takes a user message, a list of categories, and uses the AI -// mode as a zero-shot classifier. It returns an error if the classification -// result is not a valid class. -func (a *Assist) ClassifyMessage(ctx context.Context, message string, classes map[string]string) (string, error) { - category, err := a.client.ClassifyMessage(ctx, message, classes) - if err != nil { - return "", trace.Wrap(err) - } - - cleanedCategory := strings.ToLower(strings.Trim(category, ". ")) - if _, ok := classes[cleanedCategory]; ok { - return cleanedCategory, nil - } - - return "", trace.CompareFailed("classification failed, category '%s' is not a valid classes", cleanedCategory) -} - -// loadMessages loads the messages from the database. -func (c *Chat) loadMessages(ctx context.Context) error { - // existing conversation, retrieve old messages - messages, err := c.assistService.GetAssistantMessages(ctx, &assist.GetAssistantMessagesRequest{ - ConversationId: c.ConversationID, - Username: c.Username, - }) - if err != nil { - return trace.Wrap(err) - } - - // restore conversation context. - for _, msg := range messages.GetMessages() { - role := kindToRole(MessageType(msg.Type)) - if role != "" { - payload, err := formatMessagePayload(msg) - if err != nil { - return trace.Wrap(err) - } - c.chat.Insert(role, payload) - } - } - - // Mark the history as fresh. - c.potentiallyStaleHistory = false - - return nil -} - -// IsNewConversation returns true if the conversation has no messages yet. -func (c *Chat) IsNewConversation() bool { - return len(c.chat.GetMessages()) == 1 -} - -// getAssistantClient returns the OpenAI client created base on Teleport Plugin information -// or the static token configured in YAML. -func getAssistantClient(ctx context.Context, proxyClient PluginGetter, - proxySettings any, openaiCfg *openai.ClientConfig, -) (*ai.Client, error) { - apiKey, err := getOpenAITokenFromDefaultPlugin(ctx, proxyClient) - if err == nil { - return ai.NewClient(apiKey), nil - } else if !trace.IsNotFound(err) && !trace.IsNotImplemented(err) { - // We ignore 2 types of errors here. - // Unimplemented may be raised by the OSS server, - // as PluginsService does not exist there yet. - // NotFound means plugin does not exist, - // in which case we should fall back on the static token configured in YAML. - log.WithError(err).Error("Unexpected error fetching default OpenAI plugin") - } - - // If the default plugin is not configured, try to get the token from the proxy settings. - keyGetter, found := proxySettings.(interface{ GetOpenAIAPIKey() string }) - if !found { - return nil, trace.Errorf("GetOpenAIAPIKey is not implemented on %T", proxySettings) - } - - apiKey = keyGetter.GetOpenAIAPIKey() - if apiKey == "" { - return nil, trace.Errorf("OpenAI API key is not set") - } - - // Allow using the passed config if passed. - // In this case, apiKey is ignored, the one from the OpenAI config is used. - if openaiCfg != nil { - return ai.NewClientFromConfig(*openaiCfg), nil - } - return ai.NewClient(apiKey), nil -} - -// onMessageFunc is a function that is called when a message is received. -type onMessageFunc func(kind MessageType, payload []byte, createdTime time.Time) error - -// RecordMessage is used to record out-of-band messages such as hidden acknowledgements. -func (c *Chat) RecordMesssage(ctx context.Context, kind MessageType, payload string) error { - switch kind { - case MessageKindAccessRequestCreated: - protoMsg := &assist.CreateAssistantMessageRequest{ - ConversationId: c.ConversationID, - Username: c.Username, - Message: &assist.AssistantMessage{ - Type: string(MessageKindAssistantMessage), - Payload: payload, - CreatedTime: timestamppb.New(c.assist.clock.Now().UTC()), - }, - } - - if err := c.assistService.CreateAssistantMessage(ctx, protoMsg); err != nil { - return trace.Wrap(err) - } - default: - return trace.BadParameter("unsupported marker message kind: %v", kind) - } - - return nil -} - -// ProcessComplete processes the completion request and returns the number of tokens used. -func (c *Chat) ProcessComplete(ctx context.Context, onMessage onMessageFunc, userInput string, -) (*tokens.TokenCount, error) { - progressUpdates := func(update *model.AgentAction) { - payload, err := json.Marshal(update) - if err != nil { - log.WithError(err).Debugf("Failed to marshal progress update: %v", update) - return - } - - if err := onMessage(MessageKindProgressUpdate, payload, c.assist.clock.Now().UTC()); err != nil { - log.WithError(err).Debugf("Failed to send progress update: %v", update) - return - } - } - - // If data might have been inserted into the chat history, we want to - // refresh and get the latest data before querying the model. - if c.potentiallyStaleHistory { - if err := c.reloadMessages(ctx); err != nil { - return nil, trace.Wrap(err) - } - } - - // query the assistant and fetch an answer - message, tokenCount, err := c.chat.Complete(ctx, userInput, progressUpdates) - if err != nil { - return nil, trace.Wrap(err) - } - - // write the user message to persistent storage and the chat structure - c.chat.Insert(openai.ChatMessageRoleUser, userInput) - - // Do not write empty messages to the database. - if userInput != "" { - if err := c.assistService.CreateAssistantMessage(ctx, &assist.CreateAssistantMessageRequest{ - Message: &assist.AssistantMessage{ - Type: string(MessageKindUserMessage), - Payload: userInput, // TODO(jakule): Sanitize the payload - CreatedTime: timestamppb.New(c.assist.clock.Now().UTC()), - }, - ConversationId: c.ConversationID, - Username: c.Username, - }); err != nil { - return nil, trace.Wrap(err) - } - } - - switch message := message.(type) { - case *output.Message: - c.chat.Insert(openai.ChatMessageRoleAssistant, message.Content) - - // write an assistant message to persistent storage - protoMsg := &assist.CreateAssistantMessageRequest{ - ConversationId: c.ConversationID, - Username: c.Username, - Message: &assist.AssistantMessage{ - Type: string(MessageKindAssistantMessage), - Payload: message.Content, - CreatedTime: timestamppb.New(c.assist.clock.Now().UTC()), - }, - } - - if err := c.assistService.CreateAssistantMessage(ctx, protoMsg); err != nil { - return nil, trace.Wrap(err) - } - - if err := onMessage(MessageKindAssistantMessage, []byte(message.Content), c.assist.clock.Now().UTC()); err != nil { - return nil, trace.Wrap(err) - } - case *output.StreamingMessage: - var text strings.Builder - defer onMessage(MessageKindAssistantPartialFinalize, nil, c.assist.clock.Now().UTC()) - for part := range message.Parts { - text.WriteString(part) - - if err := onMessage(MessageKindAssistantPartialMessage, []byte(part), c.assist.clock.Now().UTC()); err != nil { - return nil, trace.Wrap(err) - } - } - - // write an assistant message to memory and persistent storage - textS := text.String() - c.chat.Insert(openai.ChatMessageRoleAssistant, textS) - protoMsg := &assist.CreateAssistantMessageRequest{ - ConversationId: c.ConversationID, - Username: c.Username, - Message: &assist.AssistantMessage{ - Type: string(MessageKindAssistantMessage), - Payload: textS, - CreatedTime: timestamppb.New(c.assist.clock.Now().UTC()), - }, - } - - if err := c.assistService.CreateAssistantMessage(ctx, protoMsg); err != nil { - return nil, trace.Wrap(err) - } - case *output.CompletionCommand: - payloadJson, err := json.Marshal(message) - if err != nil { - return nil, trace.Wrap(err) - } - - msg := &assist.CreateAssistantMessageRequest{ - ConversationId: c.ConversationID, - Username: c.Username, - Message: &assist.AssistantMessage{ - Type: string(MessageKindCommand), - Payload: string(payloadJson), - CreatedTime: timestamppb.New(c.assist.clock.Now().UTC()), - }, - } - - if err := c.assistService.CreateAssistantMessage(ctx, msg); err != nil { - return nil, trace.Wrap(err) - } - - if err := onMessage(MessageKindCommand, payloadJson, c.assist.clock.Now().UTC()); nil != err { - return nil, trace.Wrap(err) - } - // As we emitted a command suggestion, the user might have run it. If - // the command ran, a summary could have been inserted in the backend. - // To take this command summary into account, we note the history might - // be stale. - c.potentiallyStaleHistory = true - case *output.AccessRequest: - payloadJson, err := json.Marshal(message) - if err != nil { - return nil, trace.Wrap(err) - } - - msg := &assist.CreateAssistantMessageRequest{ - ConversationId: c.ConversationID, - Username: c.Username, - Message: &assist.AssistantMessage{ - Type: string(MessageKindAccessRequest), - Payload: string(payloadJson), - CreatedTime: timestamppb.New(c.assist.clock.Now().UTC()), - }, - } - - if err := c.assistService.CreateAssistantMessage(ctx, msg); err != nil { - return nil, trace.Wrap(err) - } - - if err := onMessage(MessageKindAccessRequest, payloadJson, c.assist.clock.Now().UTC()); nil != err { - return nil, trace.Wrap(err) - } - default: - return nil, trace.Errorf("unknown message type: %T", message) - } - - return tokenCount, nil -} - -// ProcessComplete processes a user message and returns the assistant's response. -func (c *LightweightChat) ProcessComplete(ctx context.Context, onMessage onMessageFunc, userInput string, -) (*tokens.TokenCount, error) { - progressUpdates := func(update *model.AgentAction) { - payload, err := json.Marshal(update) - if err != nil { - log.WithError(err).Debugf("Failed to marshal progress update: %v", update) - return - } - - if err := onMessage(MessageKindProgressUpdate, payload, c.assist.clock.Now().UTC()); err != nil { - log.WithError(err).Debugf("Failed to send progress update: %v", update) - return - } - } - - message, tokenCount, err := c.chat.Reply(ctx, userInput, progressUpdates) - if err != nil { - return nil, trace.Wrap(err) - } - - c.chat.Insert(openai.ChatMessageRoleUser, userInput) - - switch message := message.(type) { - case *output.Message: - c.chat.Insert(openai.ChatMessageRoleAssistant, message.Content) - if err := onMessage(MessageKindAssistantMessage, []byte(message.Content), c.assist.clock.Now().UTC()); err != nil { - return nil, trace.Wrap(err) - } - case *output.GeneratedCommand: - c.chat.Insert(openai.ChatMessageRoleAssistant, message.Command) - if err := onMessage(MessageKindCommand, []byte(message.Command), c.assist.clock.Now().UTC()); err != nil { - return nil, trace.Wrap(err) - } - case *output.StreamingMessage: - if err := func() error { - var text strings.Builder - defer onMessage(MessageKindAssistantPartialFinalize, nil, c.assist.clock.Now().UTC()) - for part := range message.Parts { - text.WriteString(part) - - if err := onMessage(MessageKindAssistantPartialMessage, []byte(part), c.assist.clock.Now().UTC()); err != nil { - return trace.Wrap(err) - } - } - c.chat.Insert(openai.ChatMessageRoleAssistant, text.String()) - return nil - }(); err != nil { - return nil, trace.Wrap(err) - } - default: - return nil, trace.Errorf("Unexpected message type: %T", message) - } - - return tokenCount, nil -} - -func getOpenAITokenFromDefaultPlugin(ctx context.Context, proxyClient PluginGetter) (string, error) { - // Try retrieving credentials from the plugin resource first - openaiPlugin, err := proxyClient.PluginsClient().GetPlugin(ctx, &pluginsv1.GetPluginRequest{ - Name: "openai-default", - WithSecrets: true, - }) - if err != nil { - return "", trail.FromGRPC(err) - } - - creds := openaiPlugin.Credentials.GetBearerToken() - if creds == nil { - return "", trace.BadParameter("malformed credentials") - } - - return creds.Token, nil -} - -// kindToRole converts a message kind to an OpenAI role. -func kindToRole(kind MessageType) string { - switch kind { - case MessageKindUserMessage: - return openai.ChatMessageRoleUser - case MessageKindAssistantMessage: - return openai.ChatMessageRoleAssistant - case MessageKindSystemMessage: - return openai.ChatMessageRoleSystem - case MessageKindCommandResultSummary: - return openai.ChatMessageRoleUser - default: - return "" - } -} - -// formatMessagePayload generates the OpemAI message payload corresponding to -// an Assist message. Most Assist message payloads can be converted directly, -// but some payloads are JSON-formatted and must be processed before being -// passed to the model. -func formatMessagePayload(message *assist.AssistantMessage) (string, error) { - switch MessageType(message.GetType()) { - case MessageKindCommandResultSummary: - var summary CommandExecSummary - err := json.Unmarshal([]byte(message.GetPayload()), &summary) - if err != nil { - return "", trace.Wrap(err) - } - return summary.String(), nil - default: - return message.GetPayload(), nil - } -} diff --git a/lib/assist/assist_test.go b/lib/assist/assist_test.go deleted file mode 100644 index 1bc678de072..00000000000 --- a/lib/assist/assist_test.go +++ /dev/null @@ -1,209 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package assist - -import ( - "context" - "errors" - "net/http/httptest" - "testing" - "time" - - "github.com/jonboulle/clockwork" - "github.com/sashabaranov/go-openai" - "github.com/stretchr/testify/require" - "google.golang.org/grpc" - "google.golang.org/protobuf/types/known/timestamppb" - - "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" - pluginsv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/plugins/v1" - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/lib/ai/model/tools" - aitest "github.com/gravitational/teleport/lib/ai/testutils" - "github.com/gravitational/teleport/lib/auth" - "github.com/gravitational/teleport/lib/modules" -) - -func TestChatComplete(t *testing.T) { - t.Parallel() - - modules.SetInsecureTestMode(true) - // Given an OpenAI server that returns a response for a chat completion request. - responses := []string{ - generateCommandResponse(), - } - - server := httptest.NewServer(aitest.GetTestHandlerFn(t, responses)) - t.Cleanup(server.Close) - - cfg := openai.DefaultConfig("secret-test-token") - cfg.BaseURL = server.URL - - // And a chat client. - ctx := context.Background() - client, err := NewClient(ctx, &mockPluginGetter{}, &apiKeyMock{}, &cfg) - require.NoError(t, err) - - // And a test auth server. - authSrv, err := auth.NewTestAuthServer(auth.TestAuthServerConfig{ - Dir: t.TempDir(), - Clock: clockwork.NewFakeClock(), - }) - require.NoError(t, err) - - // And created conversation. - toolContext := &tools.ToolContext{ - User: "bob", - } - conversationResp, err := authSrv.AuthServer.CreateAssistantConversation(ctx, &assist.CreateAssistantConversationRequest{ - Username: toolContext.User, - CreatedTime: timestamppb.Now(), - }) - require.NoError(t, err) - - // When a chat is created. - chat, err := client.NewChat(ctx, authSrv.AuthServer, toolContext, conversationResp.Id) - require.NoError(t, err) - - t.Run("new conversation is new", func(t *testing.T) { - // Then the chat is new. - require.True(t, chat.IsNewConversation()) - }) - - t.Run("the first message is the hey message", func(t *testing.T) { - // Use called to make sure that the callback is called. - called := false - // The first message is the welcome message. - _, err = chat.ProcessComplete(ctx, func(kind MessageType, payload []byte, createdTime time.Time) error { - require.Equal(t, MessageKindAssistantMessage, kind) - require.Contains(t, string(payload), "Hey, I'm Teleport") - called = true - return nil - }, "") - require.NoError(t, err) - require.True(t, called) - }) - - t.Run("command should be returned in the response", func(t *testing.T) { - called := false - // The second message is the command response. - _, err = chat.ProcessComplete(ctx, func(kind MessageType, payload []byte, createdTime time.Time) error { - if kind == MessageKindProgressUpdate { - return nil - } - require.Equal(t, MessageKindCommand, kind) - require.Equal(t, `{"command":"df -h","nodes":["localhost"]}`, string(payload)) - called = true - return nil - }, "Show free disk space on localhost") - require.NoError(t, err) - require.True(t, called) - }) - - t.Run("check what messages are stored in the backend", func(t *testing.T) { - // backend should have 3 messages: welcome message, user message, command response. - messages, err := authSrv.AuthServer.GetAssistantMessages(ctx, &assist.GetAssistantMessagesRequest{ - Username: toolContext.User, - ConversationId: conversationResp.Id, - }) - require.NoError(t, err) - require.Len(t, messages.Messages, 3) - - require.Equal(t, string(MessageKindAssistantMessage), messages.Messages[0].Type) - require.Equal(t, string(MessageKindUserMessage), messages.Messages[1].Type) - require.Equal(t, string(MessageKindCommand), messages.Messages[2].Type) - }) -} - -func TestClassifyMessage(t *testing.T) { - // Given an OpenAI server that returns a response for a chat completion request. - responses := []string{ - "troubleshooting", - "Troubleshooting", - "Troubleshooting.", - "non-existent", - } - - server := httptest.NewServer(aitest.GetTestHandlerFn(t, responses)) - t.Cleanup(server.Close) - - cfg := openai.DefaultConfig("secret-test-token") - cfg.BaseURL = server.URL - - // And a chat client. - ctx := context.Background() - client, err := NewClient(ctx, &mockPluginGetter{}, &apiKeyMock{}, &cfg) - require.NoError(t, err) - - t.Run("Valid class", func(t *testing.T) { - class, err := client.ClassifyMessage(ctx, "whatever", MessageClasses) - require.NoError(t, err) - require.Equal(t, "troubleshooting", class) - }) - - t.Run("Valid class starting with upper-case", func(t *testing.T) { - class, err := client.ClassifyMessage(ctx, "whatever", MessageClasses) - require.NoError(t, err) - require.Equal(t, "troubleshooting", class) - }) - - t.Run("Valid class starting with upper-case and ending with dot", func(t *testing.T) { - class, err := client.ClassifyMessage(ctx, "whatever", MessageClasses) - require.NoError(t, err) - require.Equal(t, "troubleshooting", class) - }) - - t.Run("Model hallucinates", func(t *testing.T) { - class, err := client.ClassifyMessage(ctx, "whatever", MessageClasses) - require.Error(t, err) - require.Empty(t, class) - }) -} - -type apiKeyMock struct{} - -// GetOpenAIAPIKey returns a mock API key. -func (m *apiKeyMock) GetOpenAIAPIKey() string { - return "123" -} - -type mockPluginGetter struct{} - -func (m *mockPluginGetter) PluginsClient() pluginsv1.PluginServiceClient { - return &mockPluginServiceClient{} -} - -type mockPluginServiceClient struct { - pluginsv1.PluginServiceClient -} - -// GetPlugin always returns an error, so the assist fallbacks to the default config. -func (m *mockPluginServiceClient) GetPlugin(_ context.Context, _ *pluginsv1.GetPluginRequest, _ ...grpc.CallOption) (*types.PluginV1, error) { - return nil, errors.New("not implemented") -} - -// generateCommandResponse generates a response for the command "df -h" on the node "localhost" -func generateCommandResponse() string { - return "```" + `json - { - "action": "Command Execution", - "action_input": "{\"command\":\"df -h\",\"nodes\":[\"localhost\"],\"labels\":[]}" - } - ` + "```" -} diff --git a/lib/assist/constants.go b/lib/assist/constants.go deleted file mode 100644 index a807980649a..00000000000 --- a/lib/assist/constants.go +++ /dev/null @@ -1,37 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package assist - -// MessageClasses contains type of assist message we expect users to send. -// When running on Cloud we attempt to classify user messages in one of those -// categories. If this succeeds, we send an event into the analytics pipeline. -// -// Keys are the category names, those are the ones reported in the event and -// generated by the model. Values are the category description. They are used to -// build the model prompt and allow to provide more context to the model to -// improve the classification. -var MessageClasses = map[string]string{ - "command execution": "the user want to execute a command on one or many servers", - "troubleshooting": "the user wants to diagnose a problem or understand an error message", - "configuration": "the user wants to generate configuration for a software which is not Teleport", - "manage resources": "the user wants to list/add/remove/edit resources connected to the Teleport cluster", - "access request": "the user requests access to one or many resources from the Teleport cluster", - "teleport setup": "the user wants help with its Teleport cluster, like setting up a new feature or knowing if something is feasible", - "other": "the user asks a question which is not IT nor Teleport-related", -} diff --git a/lib/assist/messages.go b/lib/assist/messages.go deleted file mode 100644 index 9a6faca10cd..00000000000 --- a/lib/assist/messages.go +++ /dev/null @@ -1,36 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package assist - -import ( - "fmt" -) - -// CommandExecSummary is a payload for the COMMAND_RESULT_SUMMARY message. -type CommandExecSummary struct { - ExecutionID string `json:"execution_id"` - Summary string `json:"summary"` - Command string `json:"command"` -} - -// String implements the Stringer interface and formats the message for AI -// model consumption. -func (s CommandExecSummary) String() string { - return fmt.Sprintf("Command: `%s` executed. The command output summary is: %s", s.Command, s.Summary) -} diff --git a/lib/auth/assist/assistv1/service.go b/lib/auth/assist/assistv1/service.go deleted file mode 100644 index 45c4667a44f..00000000000 --- a/lib/auth/assist/assistv1/service.go +++ /dev/null @@ -1,533 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package assistv1 - -import ( - "context" - "slices" - - "github.com/gravitational/trace" - "github.com/sirupsen/logrus" - "google.golang.org/protobuf/types/known/emptypb" - - "github.com/gravitational/teleport" - "github.com/gravitational/teleport/api/defaults" - "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/lib/ai" - embeddinglib "github.com/gravitational/teleport/lib/ai/embedding" - "github.com/gravitational/teleport/lib/authz" - "github.com/gravitational/teleport/lib/modules" - "github.com/gravitational/teleport/lib/services" -) - -const ( - // maxSearchLimit is the maximum number of search results to return. - // We have a hard cap due the simplistic design of our retriever which has quadratic complexity. - maxSearchLimit = 100 -) - -// ServiceConfig holds configuration options for -// the assist gRPC service. -type ServiceConfig struct { - Backend services.Assistant - Embeddings *ai.SimpleRetriever - Embedder embeddinglib.Embedder - Authorizer authz.Authorizer - Logger *logrus.Entry - ResourceGetter ResourceGetter -} - -// ResourceGetter represents a subset of the auth.Cache interface. -// Created to avoid circular dependencies. -type ResourceGetter interface { - GetNode(ctx context.Context, namespace, name string) (types.Server, error) - GetKubernetesCluster(ctx context.Context, name string) (types.KubeCluster, error) - GetApp(ctx context.Context, name string) (types.Application, error) - GetDatabase(ctx context.Context, name string) (types.Database, error) - GetWindowsDesktops(ctx context.Context, filter types.WindowsDesktopFilter) ([]types.WindowsDesktop, error) -} - -// Service implements the teleport.assist.v1.AssistService RPC service. -type Service struct { - assist.UnimplementedAssistServiceServer - assist.UnimplementedAssistEmbeddingServiceServer - - backend services.Assistant - embeddings *ai.SimpleRetriever - // embedder is used to embed text into a vector. - // It can be nil if the OpenAI API key is not set. - embedder embeddinglib.Embedder - authorizer authz.Authorizer - log *logrus.Entry - resourceGetter ResourceGetter -} - -// NewService returns a new assist gRPC service. -func NewService(cfg *ServiceConfig) (*Service, error) { - switch { - case cfg.Backend == nil: - return nil, trace.BadParameter("backend is required") - case cfg.Embeddings == nil: - return nil, trace.BadParameter("embeddings is required") - case cfg.Authorizer == nil: - return nil, trace.BadParameter("authorizer is required") - case cfg.ResourceGetter == nil: - return nil, trace.BadParameter("resource getter is required") - case cfg.Logger == nil: - cfg.Logger = logrus.WithField(teleport.ComponentKey, "assist.service") - } - // Embedder can be nil is the OpenAI API key is not set. - - return &Service{ - backend: cfg.Backend, - embeddings: cfg.Embeddings, - embedder: cfg.Embedder, - authorizer: cfg.Authorizer, - resourceGetter: cfg.ResourceGetter, - log: cfg.Logger, - }, nil -} - -// CreateAssistantConversation creates a new conversation entry in the backend. -func (a *Service) CreateAssistantConversation(ctx context.Context, req *assist.CreateAssistantConversationRequest) (*assist.CreateAssistantConversationResponse, error) { - authCtx, err := a.authorizer.Authorize(ctx) - if err != nil { - return nil, trace.Wrap(err) - } - - if err := authCtx.CheckAccessToKind(types.KindAssistant, types.VerbCreate); err != nil { - return nil, trace.Wrap(err) - } - - if userHasAccess(authCtx, req) { - return nil, trace.AccessDenied("user %q is not allowed to create conversation for user %q", authCtx.User.GetName(), req.Username) - } - - resp, err := a.backend.CreateAssistantConversation(ctx, req) - return resp, trace.Wrap(err) -} - -// UpdateAssistantConversationInfo updates the conversation info for a conversation. -func (a *Service) UpdateAssistantConversationInfo(ctx context.Context, req *assist.UpdateAssistantConversationInfoRequest) (*emptypb.Empty, error) { - authCtx, err := a.authorizer.Authorize(ctx) - if err != nil { - return nil, trace.Wrap(err) - } - - if err := authCtx.CheckAccessToKind(types.KindAssistant, types.VerbUpdate); err != nil { - return nil, trace.Wrap(err) - } - - if userHasAccess(authCtx, req) { - return nil, trace.AccessDenied("user %q is not allowed to update conversation for user %q", authCtx.User.GetName(), req.Username) - } - - err = a.backend.UpdateAssistantConversationInfo(ctx, req) - if err != nil { - return &emptypb.Empty{}, trace.Wrap(err) - } - - return &emptypb.Empty{}, nil -} - -// GetAssistantConversations returns all conversations started by a user. -func (a *Service) GetAssistantConversations(ctx context.Context, req *assist.GetAssistantConversationsRequest) (*assist.GetAssistantConversationsResponse, error) { - authCtx, err := a.authorizer.Authorize(ctx) - if err != nil { - return nil, trace.Wrap(err) - } - - if err := authCtx.CheckAccessToKind(types.KindAssistant, types.VerbList); err != nil { - return nil, trace.Wrap(err) - } - - if userHasAccess(authCtx, req) { - return nil, trace.AccessDenied("user %q is not allowed to list conversations for user %q", authCtx.User.GetName(), req.GetUsername()) - } - - resp, err := a.backend.GetAssistantConversations(ctx, req) - return resp, trace.Wrap(err) -} - -// DeleteAssistantConversation deletes a conversation entry and associated messages from the backend. -func (a *Service) DeleteAssistantConversation(ctx context.Context, req *assist.DeleteAssistantConversationRequest) (*emptypb.Empty, error) { - authCtx, err := a.authorizer.Authorize(ctx) - if err != nil { - return nil, trace.Wrap(err) - } - - if err := authCtx.CheckAccessToKind(types.KindAssistant, types.VerbDelete); err != nil { - return nil, trace.Wrap(err) - } - - if userHasAccess(authCtx, req) { - return nil, trace.AccessDenied("user %q is not allowed to delete conversation for user %q", authCtx.User.GetName(), req.GetUsername()) - } - - return &emptypb.Empty{}, trace.Wrap(a.backend.DeleteAssistantConversation(ctx, req)) -} - -// GetAssistantMessages returns all messages with given conversation ID. -func (a *Service) GetAssistantMessages(ctx context.Context, req *assist.GetAssistantMessagesRequest) (*assist.GetAssistantMessagesResponse, error) { - authCtx, err := a.authorizer.Authorize(ctx) - if err != nil { - return nil, trace.Wrap(err) - } - - if err := authCtx.CheckAccessToKind(types.KindAssistant, types.VerbRead); err != nil { - return nil, trace.Wrap(err) - } - - if userHasAccess(authCtx, req) { - return nil, trace.AccessDenied("user %q is not allowed to get messages for user %q", authCtx.User.GetName(), req.GetUsername()) - } - - resp, err := a.backend.GetAssistantMessages(ctx, req) - return resp, trace.Wrap(err) -} - -// CreateAssistantMessage adds the message to the backend. -func (a *Service) CreateAssistantMessage(ctx context.Context, req *assist.CreateAssistantMessageRequest) (*emptypb.Empty, error) { - authCtx, err := a.authorizer.Authorize(ctx) - if err != nil { - return nil, trace.Wrap(err) - } - - if err := authCtx.CheckAccessToKind(types.KindAssistant, types.VerbCreate); err != nil { - return nil, trace.Wrap(err) - } - - if userHasAccess(authCtx, req) { - return nil, trace.AccessDenied("user %q is not allowed to create message for user %q", authCtx.User.GetName(), req.GetUsername()) - } - - return &emptypb.Empty{}, trace.Wrap(a.backend.CreateAssistantMessage(ctx, req)) -} - -// IsAssistEnabled returns true if the assist is enabled or not on the auth level. -func (a *Service) IsAssistEnabled(ctx context.Context, _ *assist.IsAssistEnabledRequest) (*assist.IsAssistEnabledResponse, error) { - if !modules.GetModules().Features().Assist { - // If the assist feature is not enabled on the license, the assist is not enabled. - return &assist.IsAssistEnabledResponse{Enabled: false}, nil - } - - // If the embedder is not configured, the assist is not enabled as we cannot compute embeddings. - if a.embedder == nil { - return &assist.IsAssistEnabledResponse{Enabled: false}, nil - } - - authCtx, err := a.authorizer.Authorize(ctx) - if err != nil { - return nil, trace.Wrap(err) - } - - // Check if this endpoint is called by a user or Proxy. - if authz.IsLocalUser(*authCtx) { - checkErr := authCtx.Checker.CheckAccessToRule( - &services.Context{User: authCtx.User}, - defaults.Namespace, types.KindAssistant, types.VerbRead, - ) - if checkErr != nil { - return nil, trace.Wrap(err) - } - } else { - // This endpoint is called from Proxy to check if the assist is enabled. - // Proxy credentials are used instead of the user credentials. - requestedByProxy := authz.HasBuiltinRole(*authCtx, string(types.RoleProxy)) - if !requestedByProxy { - return nil, trace.AccessDenied("only proxy is allowed to call IsAssistEnabled endpoint") - } - } - - // Check if assist can use the backend. - return a.backend.IsAssistEnabled(ctx) -} - -func (a *Service) GetAssistantEmbeddings(ctx context.Context, msg *assist.GetAssistantEmbeddingsRequest) (*assist.GetAssistantEmbeddingsResponse, error) { - switch msg.Kind { - case types.KindNode, types.KindKubernetesCluster, types.KindApp, types.KindDatabase, types.KindWindowsDesktop: - default: - return nil, trace.BadParameter("resource kind %v is not supported", msg.Kind) - } - - authCtx, err := a.authorizer.Authorize(ctx) - if err != nil { - return nil, trace.Wrap(err) - } - - if err := authCtx.CheckAccessToKind(msg.Kind, types.VerbRead, types.VerbList); err != nil { - return nil, trace.Wrap(err) - } - - if a.embedder == nil { - return nil, trace.BadParameter("assist is not configured in auth server") - } - - // Call the openAI API to get the embeddings for the query. - embeddings, err := a.embedder.ComputeEmbeddings(ctx, []string{msg.Query}) - if err != nil { - return nil, trace.Wrap(err) - } - if len(embeddings) == 0 { - return nil, trace.NotFound("OpenAI embeddings returned no results") - } - - // Use default values for the id and content, as we only care about the embeddings. - queryEmbeddings := embeddinglib.NewEmbedding(msg.Kind, "", embeddings[0], [32]byte{}) - accessChecker := makeAccessChecker(ctx, a, authCtx, msg.Kind) - documents := a.embeddings.GetRelevant(queryEmbeddings, int(msg.Limit), accessChecker) - return assembleEmbeddingResponse(ctx, a, documents) -} - -// SearchUnifiedResources returns a similarity-ordered list of resources from the unified resource cache -func (a *Service) SearchUnifiedResources(ctx context.Context, msg *assist.SearchUnifiedResourcesRequest) (*assist.SearchUnifiedResourcesResponse, error) { - if a.embedder == nil { - return nil, trace.BadParameter("assist is not configured in auth server") - } - - // Call the openAI API to get the embeddings for the query. - embeddings, err := a.embedder.ComputeEmbeddings(ctx, []string{msg.Query}) - if err != nil { - return nil, trace.Wrap(err) - } - if len(embeddings) == 0 { - return nil, trace.NotFound("OpenAI embeddings returned no results") - } - - authCtx, err := a.authorizer.Authorize(ctx) - if err != nil { - return nil, trace.Wrap(err) - } - - // Use default values for the id and content, as we only care about the embeddings. - queryEmbeddings := embeddinglib.NewEmbedding("", "", embeddings[0], [32]byte{}) - limit := max(msg.Limit, maxSearchLimit) - accessChecker := makeAccessChecker(ctx, a, authCtx, msg.Kinds...) - documents := a.embeddings.GetRelevant(queryEmbeddings, int(limit), accessChecker) - return assembleSearchResponse(ctx, a, documents) -} - -// userHasAccess returns true if the user should have access to the resource. -func userHasAccess(authCtx *authz.Context, req interface{ GetUsername() string }) bool { - return !authz.IsCurrentUser(*authCtx, req.GetUsername()) && !authz.HasBuiltinRole(*authCtx, string(types.RoleAdmin)) -} - -func assembleSearchResponse(ctx context.Context, a *Service, documents []*ai.Document) (*assist.SearchUnifiedResourcesResponse, error) { - resources := make([]types.ResourceWithLabels, 0, len(documents)) - - for _, doc := range documents { - var resource types.ResourceWithLabels - var err error - - switch doc.EmbeddedKind { - case types.KindNode: - resource, err = a.resourceGetter.GetNode(ctx, defaults.Namespace, doc.GetEmbeddedID()) - case types.KindKubernetesCluster: - resource, err = a.resourceGetter.GetKubernetesCluster(ctx, doc.GetEmbeddedID()) - case types.KindApp: - resource, err = a.resourceGetter.GetApp(ctx, doc.GetEmbeddedID()) - case types.KindDatabase: - resource, err = a.resourceGetter.GetDatabase(ctx, doc.GetEmbeddedID()) - case types.KindWindowsDesktop: - desktops, err := a.resourceGetter.GetWindowsDesktops(ctx, types.WindowsDesktopFilter{ - Name: doc.GetEmbeddedID(), - }) - if err != nil { - return nil, trace.Wrap(err) - } - - for _, d := range desktops { - if d.GetName() == doc.GetEmbeddedID() { - resource = d - break - } - } - - if resource == nil { - return nil, trace.NotFound("windows desktop %q not found", doc.GetEmbeddedID()) - } - default: - return nil, trace.BadParameter("resource kind %v is not supported", doc.EmbeddedKind) - } - - if err != nil { - return nil, trace.Wrap(err) - } - - resources = append(resources, resource) - } - - paginated, err := services.MakePaginatedResources(ctx, types.KindUnifiedResource, resources, nil /* requestable map */) - if err != nil { - return nil, trace.Wrap(err) - } - - return &assist.SearchUnifiedResourcesResponse{ - Resources: paginated, - }, nil -} - -func assembleEmbeddingResponse(ctx context.Context, a *Service, documents []*ai.Document) (*assist.GetAssistantEmbeddingsResponse, error) { - protoDocs := make([]*assist.EmbeddedDocument, 0, len(documents)) - - for _, doc := range documents { - var content []byte - - switch doc.EmbeddedKind { - case types.KindNode: - node, err := a.resourceGetter.GetNode(ctx, defaults.Namespace, doc.GetEmbeddedID()) - if err != nil { - return nil, trace.Wrap(err) - } - - content, err = embeddinglib.SerializeNode(node) - if err != nil { - return nil, trace.Wrap(err) - } - case types.KindKubernetesCluster: - cluster, err := a.resourceGetter.GetKubernetesCluster(ctx, doc.GetEmbeddedID()) - if err != nil { - return nil, trace.Wrap(err) - } - - content, err = embeddinglib.SerializeKubeCluster(cluster) - if err != nil { - return nil, trace.Wrap(err) - } - case types.KindApp: - app, err := a.resourceGetter.GetApp(ctx, doc.GetEmbeddedID()) - if err != nil { - return nil, trace.Wrap(err) - } - - content, err = embeddinglib.SerializeApp(app) - if err != nil { - return nil, trace.Wrap(err) - } - case types.KindDatabase: - db, err := a.resourceGetter.GetDatabase(ctx, doc.GetEmbeddedID()) - if err != nil { - return nil, trace.Wrap(err) - } - - content, err = embeddinglib.SerializeDatabase(db) - if err != nil { - return nil, trace.Wrap(err) - } - case types.KindWindowsDesktop: - desktops, err := a.resourceGetter.GetWindowsDesktops(ctx, types.WindowsDesktopFilter{ - Name: doc.GetEmbeddedID(), - }) - if err != nil { - return nil, trace.Wrap(err) - } - - var desktop types.WindowsDesktop - for _, d := range desktops { - if d.GetName() == doc.GetEmbeddedID() { - desktop = d - break - } - } - - if desktop == nil { - return nil, trace.NotFound("windows desktop %q not found", doc.GetEmbeddedID()) - } - - content, err = embeddinglib.SerializeWindowsDesktop(desktop) - if err != nil { - return nil, trace.Wrap(err) - } - } - - protoDocs = append(protoDocs, &assist.EmbeddedDocument{ - Id: doc.GetEmbeddedID(), - Content: string(content), - SimilarityScore: float32(doc.SimilarityScore), - }) - } - - return &assist.GetAssistantEmbeddingsResponse{ - Embeddings: protoDocs, - }, nil -} - -func makeAccessChecker(ctx context.Context, a *Service, authCtx *authz.Context, kinds ...string) func(id string, embedding *embeddinglib.Embedding) bool { - return func(id string, embedding *embeddinglib.Embedding) bool { - if !slices.Contains(kinds, embedding.EmbeddedKind) && len(kinds) > 0 { - return false - } - - var resource services.AccessCheckable - var err error - - switch embedding.EmbeddedKind { - case types.KindNode: - resource, err = a.resourceGetter.GetNode(ctx, defaults.Namespace, embedding.GetEmbeddedID()) - if err != nil { - a.log.Tracef("failed to get node %q: %v", embedding.GetName(), err) - return false - } - case types.KindKubernetesCluster: - resource, err = a.resourceGetter.GetKubernetesCluster(ctx, embedding.GetEmbeddedID()) - if err != nil { - a.log.Tracef("failed to get kube cluster %q: %v", embedding.GetName(), err) - return false - } - case types.KindApp: - resource, err = a.resourceGetter.GetApp(ctx, embedding.GetEmbeddedID()) - if err != nil { - a.log.Tracef("failed to get app %q: %v", embedding.GetName(), err) - return false - } - case types.KindDatabase: - resource, err = a.resourceGetter.GetDatabase(ctx, embedding.GetEmbeddedID()) - if err != nil { - a.log.Tracef("failed to get database %q: %v", embedding.GetName(), err) - return false - } - case types.KindWindowsDesktop: - desktops, err := a.resourceGetter.GetWindowsDesktops(ctx, types.WindowsDesktopFilter{ - Name: embedding.GetEmbeddedID(), - }) - if err != nil { - a.log.Tracef("failed to get windows desktop %q: %v", embedding.GetName(), err) - return false - } - - for _, d := range desktops { - if d.GetName() == embedding.GetEmbeddedID() { - resource = d - break - } - } - - if resource == nil { - a.log.Tracef("failed to find windows desktop %q: %v", embedding.GetName(), err) - return false - } - default: - a.log.Tracef("resource kind %v is not supported", embedding.EmbeddedKind) - return false - } - - return authCtx.Checker.CheckAccess(resource, services.AccessState{MFAVerified: true}) == nil - } -} diff --git a/lib/auth/assist/assistv1/test/service_test.go b/lib/auth/assist/assistv1/test/service_test.go deleted file mode 100644 index 8e40366b810..00000000000 --- a/lib/auth/assist/assistv1/test/service_test.go +++ /dev/null @@ -1,525 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package assistv1_test - -import ( - "context" - "testing" - "time" - - "github.com/gravitational/trace" - "github.com/jonboulle/clockwork" - log "github.com/sirupsen/logrus" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - "google.golang.org/protobuf/types/known/timestamppb" - - "github.com/gravitational/teleport" - assistpb "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/api/utils/retryutils" - "github.com/gravitational/teleport/lib/ai" - "github.com/gravitational/teleport/lib/assist" - "github.com/gravitational/teleport/lib/auth/assist/assistv1" - "github.com/gravitational/teleport/lib/authz" - "github.com/gravitational/teleport/lib/backend/memory" - "github.com/gravitational/teleport/lib/defaults" - "github.com/gravitational/teleport/lib/services" - "github.com/gravitational/teleport/lib/services/local" - "github.com/gravitational/teleport/lib/tlsca" -) - -const ( - defaultUser = "test-user" - noAccessUser = "user-no-access" -) - -func TestService_CreateAssistantConversation(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - username string - req *assistpb.CreateAssistantConversationRequest - wantErr assert.ErrorAssertionFunc - assertResponse func(t *testing.T, resp *assistpb.CreateAssistantConversationResponse) - }{ - { - name: "success", - username: defaultUser, - req: &assistpb.CreateAssistantConversationRequest{ - Username: defaultUser, - CreatedTime: timestamppb.Now(), - }, - wantErr: assert.NoError, - assertResponse: func(t *testing.T, resp *assistpb.CreateAssistantConversationResponse) { - require.NotEmpty(t, resp.GetId()) - }, - }, - { - name: "access denies - RBAC", - username: noAccessUser, - req: &assistpb.CreateAssistantConversationRequest{ - Username: noAccessUser, - CreatedTime: timestamppb.Now(), - }, - wantErr: assert.Error, - }, - { - name: "access denied - different user", - username: defaultUser, - req: &assistpb.CreateAssistantConversationRequest{ - Username: noAccessUser, - CreatedTime: timestamppb.Now(), - }, - wantErr: assert.Error, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - ctxs, svc := initSvc(t) - - got, err := svc.CreateAssistantConversation(ctxs[tt.username], tt.req) - tt.wantErr(t, err) - - if tt.assertResponse != nil { - tt.assertResponse(t, got) - } - }) - } -} - -func TestService_GetAssistantConversations(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - username string - req *assistpb.GetAssistantConversationsRequest - wantErr assert.ErrorAssertionFunc - assertResponse func(t *testing.T, resp *assistpb.CreateAssistantConversationResponse) - }{ - { - name: "success", - username: defaultUser, - req: &assistpb.GetAssistantConversationsRequest{ - Username: defaultUser, - }, - wantErr: assert.NoError, - }, - { - name: "access denies - RBAC", - username: noAccessUser, - req: &assistpb.GetAssistantConversationsRequest{ - Username: noAccessUser, - }, - wantErr: assert.Error, - }, - { - name: "access denied - different user", - username: defaultUser, - req: &assistpb.GetAssistantConversationsRequest{ - Username: noAccessUser, - }, - wantErr: assert.Error, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - ctxs, svc := initSvc(t) - - _, err := svc.GetAssistantConversations(ctxs[tt.username], tt.req) - tt.wantErr(t, err) - }) - } -} - -func TestService_DeleteAssistantConversations(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - username string - req *assistpb.DeleteAssistantConversationRequest - wantConvErr assert.ErrorAssertionFunc - wantErr assert.ErrorAssertionFunc - }{ - { - name: "success", - username: defaultUser, - req: &assistpb.DeleteAssistantConversationRequest{ - Username: defaultUser, - }, - wantConvErr: assert.NoError, - wantErr: assert.NoError, - }, - { - name: "access denies - RBAC", - username: noAccessUser, - req: &assistpb.DeleteAssistantConversationRequest{ - Username: noAccessUser, - }, - wantConvErr: assert.Error, - wantErr: assert.Error, - }, - { - name: "access denied - different user", - username: defaultUser, - req: &assistpb.DeleteAssistantConversationRequest{ - Username: noAccessUser, - }, - wantConvErr: assert.NoError, - wantErr: assert.Error, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - ctxs, svc := initSvc(t) - - // Create a conversation that we can remove, so we don't hit "conversation doesn't exist" error - convMsg, err := svc.CreateAssistantConversation(ctxs[tt.username], &assistpb.CreateAssistantConversationRequest{ - Username: tt.username, - CreatedTime: timestamppb.Now(), - }) - tt.wantConvErr(t, err) - - conversationID := convMsg.GetId() - - tt.req.ConversationId = conversationID - - _, err = svc.DeleteAssistantConversation(ctxs[tt.username], tt.req) - tt.wantErr(t, err) - }) - } -} - -func TestService_InsertAssistantMessage(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - username string - req *assistpb.CreateAssistantMessageRequest - wantConvErr assert.ErrorAssertionFunc - wantErr assert.ErrorAssertionFunc - }{ - { - name: "success", - username: defaultUser, - req: &assistpb.CreateAssistantMessageRequest{ - Username: defaultUser, - Message: &assistpb.AssistantMessage{ - Type: string(assist.MessageKindAssistantMessage), - CreatedTime: timestamppb.Now(), - Payload: "Blah", - }, - }, - wantConvErr: assert.NoError, - wantErr: assert.NoError, - }, - { - name: "access denies - RBAC", - username: noAccessUser, - req: &assistpb.CreateAssistantMessageRequest{ - Username: noAccessUser, - }, - wantConvErr: assert.Error, - wantErr: assert.Error, - }, - { - name: "access denied - different user", - username: defaultUser, - req: &assistpb.CreateAssistantMessageRequest{ - Username: noAccessUser, - }, - wantConvErr: assert.NoError, - wantErr: assert.Error, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - ctxs, svc := initSvc(t) - - // Create a conversation that we can remove, so we don't hit "conversation doesn't exist" error - convMsg, err := svc.CreateAssistantConversation(ctxs[tt.username], &assistpb.CreateAssistantConversationRequest{ - Username: tt.username, - CreatedTime: timestamppb.Now(), - }) - tt.wantConvErr(t, err) - - conversationID := convMsg.GetId() - - tt.req.ConversationId = conversationID - - _, err = svc.CreateAssistantMessage(ctxs[tt.username], tt.req) - tt.wantErr(t, err) - }) - } -} - -func TestService_SearchUnifiedResources(t *testing.T) { - t.Parallel() - - tests := []struct { - username string - req *assistpb.SearchUnifiedResourcesRequest - returnedLen int - }{ - { - username: defaultUser, - req: &assistpb.SearchUnifiedResourcesRequest{ - Kinds: []string{types.KindNode}, - }, - returnedLen: 2, - }, - { - username: noAccessUser, - req: &assistpb.SearchUnifiedResourcesRequest{ - Kinds: []string{types.KindNode}, - }, - returnedLen: 0, - }, - } - - for _, tt := range tests { - t.Run(tt.username, func(t *testing.T) { - ctxs, svc := initSvc(t) - require.Eventually(t, func() bool { - resp, err := svc.SearchUnifiedResources(ctxs[tt.username], tt.req) - require.NoError(t, err) - return tt.returnedLen == len(resp.GetResources()) - }, 5*time.Second, 100*time.Millisecond) - }) - } -} - -type testClient struct { - services.ClusterConfiguration - services.Trust - services.RoleGetter - services.UserGetter -} - -func initSvc(t *testing.T) (map[string]context.Context, *assistv1.Service) { - ctx := context.Background() - backend, err := memory.New(memory.Config{}) - require.NoError(t, err) - - clusterConfigSvc, err := local.NewClusterConfigurationService(backend) - require.NoError(t, err) - trustSvc := local.NewCAService(backend) - roleSvc := local.NewAccessService(backend) - userSvc := local.NewTestIdentityService(backend) - presenceSvc := local.NewPresenceService(backend) - - _, err = clusterConfigSvc.UpsertAuthPreference(ctx, types.DefaultAuthPreference()) - require.NoError(t, err) - require.NoError(t, clusterConfigSvc.SetClusterAuditConfig(ctx, types.DefaultClusterAuditConfig())) - _, err = clusterConfigSvc.UpsertClusterNetworkingConfig(ctx, types.DefaultClusterNetworkingConfig()) - require.NoError(t, err) - _, err = clusterConfigSvc.UpsertSessionRecordingConfig(ctx, types.DefaultSessionRecordingConfig()) - require.NoError(t, err) - - accessPoint := &testClient{ - ClusterConfiguration: clusterConfigSvc, - Trust: trustSvc, - RoleGetter: roleSvc, - UserGetter: userSvc, - } - - n1, err := types.NewServer("node-1", types.KindNode, types.ServerSpecV2{}) - require.NoError(t, err) - n2, err := types.NewServer("node-2", types.KindNode, types.ServerSpecV2{}) - require.NoError(t, err) - _, err = presenceSvc.UpsertNode(ctx, n1) - require.NoError(t, err) - _, err = presenceSvc.UpsertNode(ctx, n2) - require.NoError(t, err) - - accesslistSvc, err := local.NewAccessListService(backend, clockwork.NewFakeClock()) - require.NoError(t, err) - - accessService := local.NewAccessService(backend) - eventService := local.NewEventsService(backend) - lockWatcher, err := services.NewLockWatcher(ctx, services.LockWatcherConfig{ - ResourceWatcherConfig: services.ResourceWatcherConfig{ - Client: eventService, - Component: "test", - }, - LockGetter: accessService, - }) - require.NoError(t, err) - - authorizer, err := authz.NewAuthorizer(authz.AuthorizerOpts{ - ClusterName: "test-cluster", - AccessPoint: accessPoint, - LockWatcher: lockWatcher, - }) - require.NoError(t, err) - - roles := map[string]types.Role{} - - role, err := types.NewRole("allow-rules", types.RoleSpecV6{ - Allow: types.RoleConditions{ - Namespaces: []string{}, - NodeLabels: types.Labels{types.Wildcard: []string{types.Wildcard}}, - Rules: []types.Rule{ - { - Resources: []string{types.KindAssistant}, - Verbs: []string{types.VerbList, types.VerbRead, types.VerbUpdate, types.VerbCreate, types.VerbDelete}, - }, - { - Resources: []string{types.KindNode}, - Verbs: []string{types.VerbList, types.VerbRead}, - }, - }, - }, - }) - require.NoError(t, err) - - roles[defaultUser] = role - - roleNoAccess, err := types.NewRole("no-rules", types.RoleSpecV6{ - Allow: types.RoleConditions{}, - }) - require.NoError(t, err) - roles["user-no-access"] = roleNoAccess - - ctxs := make(map[string]context.Context, len(roles)) - for username, role := range roles { - role, err = roleSvc.CreateRole(ctx, role) - require.NoError(t, err) - - user, err := types.NewUser(username) - user.AddRole(role.GetName()) - require.NoError(t, err) - - user, err = userSvc.CreateUser(ctx, user) - require.NoError(t, err) - - ctx = authz.ContextWithUser(ctx, authz.LocalUser{ - Username: user.GetName(), - Identity: tlsca.Identity{ - Username: user.GetName(), - Groups: []string{role.GetName()}, - }, - }) - ctxs[user.GetName()] = ctx - } - - embedder := ai.MockEmbedder{ - TimesCalled: make(map[string]int), - } - - embeddings := &ai.SimpleRetriever{} - embeddingSrv := local.NewEmbeddingsService(backend) - svc, err := assistv1.NewService(&assistv1.ServiceConfig{ - Backend: local.NewAssistService(backend), - Authorizer: authorizer, - Embeddings: embeddings, - ResourceGetter: &resourceGetterAllImpl{ - PresenceService: presenceSvc, - AccessLists: accesslistSvc, - }, - Embedder: &embedder, - }) - require.NoError(t, err) - - unifiedResourcesCache, err := services.NewUnifiedResourceCache(ctx, services.UnifiedResourceCacheConfig{ - ResourceWatcherConfig: services.ResourceWatcherConfig{ - QueueSize: defaults.UnifiedResourcesQueueSize, - Component: teleport.ComponentUnifiedResource, - Client: eventService, - MaxStaleness: time.Second, - }, - ResourceGetter: &resourceGetterAllImpl{ - PresenceService: presenceSvc, - AccessLists: accesslistSvc, - }, - }) - require.NoError(t, err) - - log.Debugf("Starting embedding watcher") - embeddingProcessor := ai.NewEmbeddingProcessor(&ai.EmbeddingProcessorConfig{ - AIClient: &embedder, - EmbeddingsRetriever: embeddings, - EmbeddingSrv: embeddingSrv, - NodeSrv: unifiedResourcesCache, - Jitter: retryutils.NewFullJitter(), - Log: log.NewEntry(log.StandardLogger()), - }) - log.Debugf("Starting embedding processor") - - embeddingProcessorCtx, embeddingProcessorCancel := context.WithCancel(context.Background()) - go embeddingProcessor.Run(embeddingProcessorCtx, time.Millisecond*100, time.Millisecond*100) - t.Cleanup(embeddingProcessorCancel) - return ctxs, svc -} - -type resourceGetterAllImpl struct { - *local.PresenceService - services.AccessLists -} - -func (g *resourceGetterAllImpl) GetKubernetesCluster(ctx context.Context, name string) (types.KubeCluster, error) { - kubeServers, err := g.PresenceService.GetKubernetesServers(ctx) - if err != nil { - return nil, err - } - - for _, kubeServer := range kubeServers { - if kubeServer.GetName() == name { - return kubeServer.GetCluster(), nil - } - } - - return nil, trace.NotFound("cluster not found") -} - -func (g *resourceGetterAllImpl) GetApp(ctx context.Context, name string) (types.Application, error) { - return nil, nil -} - -func (g *resourceGetterAllImpl) GetDatabase(ctx context.Context, name string) (types.Database, error) { - return nil, nil -} - -func (g *resourceGetterAllImpl) GetWindowsDesktops(ctx context.Context, _ types.WindowsDesktopFilter) ([]types.WindowsDesktop, error) { - return nil, nil -} - -func (m *resourceGetterAllImpl) GetDatabaseServers(_ context.Context, _ string, _ ...services.MarshalOption) ([]types.DatabaseServer, error) { - return nil, nil -} - -func (m *resourceGetterAllImpl) GetKubernetesServers(_ context.Context) ([]types.KubeServer, error) { - return nil, nil -} - -func (m *resourceGetterAllImpl) GetApplicationServers(_ context.Context, _ string) ([]types.AppServer, error) { - return nil, nil -} - -func (m *resourceGetterAllImpl) ListSAMLIdPServiceProviders(_ context.Context, _ int, _ string) ([]types.SAMLIdPServiceProvider, string, error) { - return nil, "", nil -} diff --git a/lib/auth/auth.go b/lib/auth/auth.go index d6b9e71a21f..78c9c42708a 100644 --- a/lib/auth/auth.go +++ b/lib/auth/auth.go @@ -66,7 +66,6 @@ import ( "github.com/gravitational/teleport/api/client/secreport" "github.com/gravitational/teleport/api/constants" apidefaults "github.com/gravitational/teleport/api/defaults" - "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" devicepb "github.com/gravitational/teleport/api/gen/proto/go/teleport/devicetrust/v1" headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1" mfav1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/mfa/v1" @@ -80,8 +79,6 @@ import ( "github.com/gravitational/teleport/api/utils/keys" "github.com/gravitational/teleport/api/utils/retryutils" apisshutils "github.com/gravitational/teleport/api/utils/sshutils" - "github.com/gravitational/teleport/lib/ai" - "github.com/gravitational/teleport/lib/ai/embedding" "github.com/gravitational/teleport/lib/auth/authclient" "github.com/gravitational/teleport/lib/auth/keystore" "github.com/gravitational/teleport/lib/auth/native" @@ -212,9 +209,6 @@ func NewServer(cfg *InitConfig, opts ...ServerOption) (*Server, error) { if cfg.Status == nil { cfg.Status = local.NewStatusService(cfg.Backend) } - if cfg.Assist == nil { - cfg.Assist = local.NewAssistService(cfg.Backend) - } if cfg.Events == nil { cfg.Events = local.NewEventsService(cfg.Backend) } @@ -312,9 +306,6 @@ func NewServer(cfg *InitConfig, opts ...ServerOption) (*Server, error) { return nil, trace.Wrap(err) } } - if cfg.Embeddings == nil { - cfg.Embeddings = local.NewEmbeddingsService(cfg.Backend) - } if cfg.UserPreferences == nil { cfg.UserPreferences = local.NewUserPreferencesService(cfg.Backend) } @@ -407,7 +398,6 @@ func NewServer(cfg *InitConfig, opts ...ServerOption) (*Server, error) { ConnectionsDiagnostic: cfg.ConnectionsDiagnostic, Integrations: cfg.Integrations, DiscoveryConfigs: cfg.DiscoveryConfigs, - Embeddings: cfg.Embeddings, Okta: cfg.Okta, AccessLists: cfg.AccessLists, DatabaseObjectImportRules: cfg.DatabaseObjectImportRules, @@ -416,7 +406,6 @@ func NewServer(cfg *InitConfig, opts ...ServerOption) (*Server, error) { UserLoginStates: cfg.UserLoginState, StatusInternal: cfg.Status, UsageReporter: cfg.UsageReporter, - Assistant: cfg.Assist, UserPreferences: cfg.UserPreferences, PluginData: cfg.PluginData, KubeWaitingContainer: cfg.KubeWaitingContainers, @@ -445,8 +434,6 @@ func NewServer(cfg *InitConfig, opts ...ServerOption) (*Server, error) { fips: cfg.FIPS, loadAllCAs: cfg.LoadAllCAs, httpClientForAWSSTS: cfg.HTTPClientForAWSSTS, - embeddingsRetriever: cfg.EmbeddingRetriever, - embedder: cfg.EmbeddingClient, accessMonitoringEnabled: cfg.AccessMonitoringEnabled, } as.inventory = inventory.NewController(&as, services, @@ -588,8 +575,6 @@ type Services struct { services.DatabaseObjectImportRules services.DatabaseObjects services.UserLoginStates - services.Assistant - services.Embeddings services.UserPreferences services.PluginData services.SCIM @@ -974,12 +959,6 @@ type Server struct { // STS requests. httpClientForAWSSTS utils.HTTPDoClient - // embeddingRetriever is a retriever used to retrieve embeddings from the backend. - embeddingsRetriever *ai.SimpleRetriever - - // embedder is an embedder client used to generate embeddings. - embedder embedding.Embedder - // accessMonitoringEnabled is a flag that indicates whether access monitoring is enabled. accessMonitoringEnabled bool @@ -6855,39 +6834,6 @@ func (a *Server) UpsertHeadlessAuthenticationStub(ctx context.Context, username return trace.Wrap(err) } -// GetAssistantMessages returns all messages with given conversation ID. -func (a *Server) GetAssistantMessages(ctx context.Context, req *assist.GetAssistantMessagesRequest) (*assist.GetAssistantMessagesResponse, error) { - resp, err := a.Services.GetAssistantMessages(ctx, req) - return resp, trace.Wrap(err) -} - -// CreateAssistantMessage adds the message to the backend. -func (a *Server) CreateAssistantMessage(ctx context.Context, msg *assist.CreateAssistantMessageRequest) error { - return trace.Wrap(a.Services.CreateAssistantMessage(ctx, msg)) -} - -// UpdateAssistantConversationInfo stores the given conversation title in the backend. -func (a *Server) UpdateAssistantConversationInfo(ctx context.Context, msg *assist.UpdateAssistantConversationInfoRequest) error { - return trace.Wrap(a.Services.UpdateAssistantConversationInfo(ctx, msg)) -} - -// CreateAssistantConversation creates a new conversation entry in the backend. -func (a *Server) CreateAssistantConversation(ctx context.Context, req *assist.CreateAssistantConversationRequest) (*assist.CreateAssistantConversationResponse, error) { - resp, err := a.Services.CreateAssistantConversation(ctx, req) - return resp, trace.Wrap(err) -} - -// GetAssistantConversations returns all conversations started by a user. -func (a *Server) GetAssistantConversations(ctx context.Context, request *assist.GetAssistantConversationsRequest) (*assist.GetAssistantConversationsResponse, error) { - resp, err := a.Services.GetAssistantConversations(ctx, request) - return resp, trace.Wrap(err) -} - -// DeleteAssistantConversation deletes a conversation from the backend. -func (a *Server) DeleteAssistantConversation(ctx context.Context, request *assist.DeleteAssistantConversationRequest) error { - return trace.Wrap(a.Services.DeleteAssistantConversation(ctx, request)) -} - // CompareAndSwapHeadlessAuthentication performs a compare // and swap replacement on a headless authentication resource. func (a *Server) CompareAndSwapHeadlessAuthentication(ctx context.Context, old, new *types.HeadlessAuthentication) (*types.HeadlessAuthentication, error) { diff --git a/lib/auth/authclient/clt.go b/lib/auth/authclient/clt.go index 99c5fa83992..055fe08d3c9 100644 --- a/lib/auth/authclient/clt.go +++ b/lib/auth/authclient/clt.go @@ -37,7 +37,6 @@ import ( "github.com/gravitational/teleport/api/client/proto" "github.com/gravitational/teleport/api/client/secreport" apidefaults "github.com/gravitational/teleport/api/defaults" - assistpb "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" clusterconfigpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/clusterconfig/v1" dbobjectimportrulev1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/dbobjectimportrule/v1" devicepb "github.com/gravitational/teleport/api/gen/proto/go/teleport/devicetrust/v1" @@ -1422,7 +1421,6 @@ type ClientI interface { services.WindowsDesktops services.SAMLIdPServiceProviders services.UserGroups - services.Assistant WebService services.Status services.ClusterConfiguration @@ -1452,9 +1450,6 @@ type ClientI interface { // "not implemented" errors (as per the default gRPC behavior). LoginRuleClient() loginrulepb.LoginRuleServiceClient - // EmbeddingClient returns a client to the Embedding gRPC service. - EmbeddingClient() assistpb.AssistEmbeddingServiceClient - // AccessGraphClient returns a client to the Access Graph gRPC service. AccessGraphClient() accessgraphv1.AccessGraphServiceClient diff --git a/lib/auth/grpcserver.go b/lib/auth/grpcserver.go index 88016750733..2042723debb 100644 --- a/lib/auth/grpcserver.go +++ b/lib/auth/grpcserver.go @@ -49,7 +49,6 @@ import ( "github.com/gravitational/teleport/api/client" authpb "github.com/gravitational/teleport/api/client/proto" "github.com/gravitational/teleport/api/constants" - "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" accessmonitoringrules "github.com/gravitational/teleport/api/gen/proto/go/teleport/accessmonitoringrules/v1" auditlogpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/auditlog/v1" clusterconfigpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/clusterconfig/v1" @@ -77,7 +76,6 @@ import ( "github.com/gravitational/teleport/api/types/installers" "github.com/gravitational/teleport/api/types/wrappers" "github.com/gravitational/teleport/lib/accessmonitoringrules/accessmonitoringrulesv1" - "github.com/gravitational/teleport/lib/auth/assist/assistv1" "github.com/gravitational/teleport/lib/auth/authclient" "github.com/gravitational/teleport/lib/auth/clusterconfig/clusterconfigv1" "github.com/gravitational/teleport/lib/auth/crownjewel/crownjewelv1" @@ -5237,20 +5235,6 @@ func NewGRPCServer(cfg GRPCServerConfig) (*GRPCServer, error) { } trustpb.RegisterTrustServiceServer(server, trust) - // Initialize and register the assist service. - assistSrv, err := assistv1.NewService(&assistv1.ServiceConfig{ - Backend: cfg.AuthServer.Services, - Embeddings: cfg.AuthServer.embeddingsRetriever, - Embedder: cfg.AuthServer.embedder, - Authorizer: cfg.Authorizer, - ResourceGetter: cfg.AuthServer, - }) - if err != nil { - return nil, trace.Wrap(err) - } - assist.RegisterAssistServiceServer(server, assistSrv) - assist.RegisterAssistEmbeddingServiceServer(server, assistSrv) - // create server with no-op role to pass to JoinService server serverWithNopRole, err := serverWithNopRole(cfg) if err != nil { diff --git a/lib/auth/helpers.go b/lib/auth/helpers.go index 499479ef800..ea2cd8f087e 100644 --- a/lib/auth/helpers.go +++ b/lib/auth/helpers.go @@ -42,8 +42,6 @@ import ( "github.com/gravitational/teleport/api/constants" "github.com/gravitational/teleport/api/types" apiutils "github.com/gravitational/teleport/api/utils" - "github.com/gravitational/teleport/lib/ai" - "github.com/gravitational/teleport/lib/ai/embedding" "github.com/gravitational/teleport/lib/auth/accesspoint" "github.com/gravitational/teleport/lib/auth/authclient" "github.com/gravitational/teleport/lib/auth/keystore" @@ -89,8 +87,6 @@ type TestAuthServerConfig struct { TraceClient otlptrace.Client // AuthPreferenceSpec is custom initial AuthPreference spec for the test. AuthPreferenceSpec *types.AuthPreferenceSpecV2 - // Embedder is required to enable the assist in the auth server. - Embedder embedding.Embedder // CacheEnabled enables the primary auth server cache. CacheEnabled bool // RunWhileLockedRetryInterval is the interval to retry the run while locked @@ -118,9 +114,6 @@ func (cfg *TestAuthServerConfig) CheckAndSetDefaults() error { SecondFactor: constants.SecondFactorOff, } } - if cfg.Embedder == nil { - cfg.Embedder = &noopEmbedder{} - } return nil } @@ -210,14 +203,6 @@ func WithClock(clock clockwork.Clock) ServerOption { } } -// WithEmbedder is a functional server option that sets the server's embedder. -func WithEmbedder(embedder embedding.Embedder) ServerOption { - return func(s *Server) error { - s.embedder = embedder - return nil - } -} - // TestAuthServer is auth server using local filesystem backend // and test certificate authority key generation that speeds up // keygen by using the same private key @@ -305,12 +290,10 @@ func NewTestAuthServer(cfg TestAuthServerConfig) (*TestAuthServer, error) { RSAKeyPairSource: authority.New().GenerateKeyPair, }, }, - EmbeddingRetriever: ai.NewSimpleRetriever(), - HostUUID: uuid.New().String(), - AccessLists: accessLists, + HostUUID: uuid.New().String(), + AccessLists: accessLists, }, WithClock(cfg.Clock), - WithEmbedder(cfg.Embedder), ) if err != nil { return nil, trace.Wrap(err) @@ -1338,13 +1321,6 @@ func CreateUserAndRoleWithoutRoles(clt clt, username string, allowedLogins []str return created, upsertedRole, nil } -// noopEmbedder is a no op implementation of the Embedder interface. -type noopEmbedder struct{} - -func (n noopEmbedder) ComputeEmbeddings(_ context.Context, _ []string) ([]embedding.Vector64, error) { - return []embedding.Vector64{}, nil -} - // flushClt is the set of methods expected by the flushCache helper. type flushClt interface { // GetNamespace returns namespace by name diff --git a/lib/auth/init.go b/lib/auth/init.go index ea1c27b2503..0149f653ecb 100644 --- a/lib/auth/init.go +++ b/lib/auth/init.go @@ -47,8 +47,6 @@ import ( "github.com/gravitational/teleport/api/types" apievents "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/lib" - "github.com/gravitational/teleport/lib/ai" - "github.com/gravitational/teleport/lib/ai/embedding" "github.com/gravitational/teleport/lib/auth/dbobjectimportrule/dbobjectimportrulev1" "github.com/gravitational/teleport/lib/auth/keystore" "github.com/gravitational/teleport/lib/auth/machineid/machineidv1" @@ -160,9 +158,6 @@ type InitConfig struct { // Status is a service that manages cluster status info. Status services.StatusInternal - // Assist is a service that implements the Teleport Assist functionality. - Assist services.Assistant - // UserPreferences is a service that manages user preferences. UserPreferences services.UserPreferences @@ -221,9 +216,6 @@ type InitConfig struct { // DiscoveryConfigs is a service that manages DiscoveryConfigs. DiscoveryConfigs services.DiscoveryConfigs - // Embeddings is a service that manages Embeddings - Embeddings services.Embeddings - // SessionTrackerService is a service that manages trackers for all active sessions. SessionTrackerService services.SessionTrackerService @@ -277,12 +269,6 @@ type InitConfig struct { // STS requests. Used in test. HTTPClientForAWSSTS utils.HTTPDoClient - // EmbeddingRetriever is a retriever for embeddings. - EmbeddingRetriever *ai.SimpleRetriever - - // EmbeddingClient is a client that allows generating embeddings. - EmbeddingClient embedding.Embedder - // Tracer used to create spans. Tracer oteltrace.Tracer diff --git a/lib/auth/userpreferences/userpreferencesv1/service_test.go b/lib/auth/userpreferences/userpreferencesv1/service_test.go index 52d7f4253e3..e7eb0061aaf 100644 --- a/lib/auth/userpreferences/userpreferencesv1/service_test.go +++ b/lib/auth/userpreferences/userpreferencesv1/service_test.go @@ -55,10 +55,6 @@ func TestService_GetUserPreferences(t *testing.T) { req: &userpreferencesv1.GetUserPreferencesRequest{}, want: &userpreferencesv1.GetUserPreferencesResponse{ Preferences: &userpreferencesv1.UserPreferences{ - Assist: &userpreferencesv1.AssistUserPreferences{ - PreferredLogins: []string{}, - ViewMode: userpreferencesv1.AssistViewMode_ASSIST_VIEW_MODE_DOCKED, - }, Theme: userpreferencesv1.Theme_THEME_UNSPECIFIED, UnifiedResourcePreferences: &userpreferencesv1.UnifiedResourcePreferences{ DefaultTab: userpreferencesv1.DefaultTab_DEFAULT_TAB_ALL, @@ -104,10 +100,6 @@ func TestService_UpsertUserPreferences(t *testing.T) { t.Parallel() defaultPreferences := &userpreferencesv1.UserPreferences{ - Assist: &userpreferencesv1.AssistUserPreferences{ - PreferredLogins: []string{}, - ViewMode: userpreferencesv1.AssistViewMode_ASSIST_VIEW_MODE_DOCKED, - }, Theme: userpreferencesv1.Theme_THEME_LIGHT, Onboard: &userpreferencesv1.OnboardUserPreferences{ PreferredResources: []userpreferencesv1.Resource{}, diff --git a/lib/config/configuration.go b/lib/config/configuration.go index 0be56381c16..641a516ea95 100644 --- a/lib/config/configuration.go +++ b/lib/config/configuration.go @@ -960,19 +960,6 @@ func applyAuthConfig(fc *FileConfig, cfg *servicecfg.Config) error { cfg.Auth.Preference.SetDisconnectExpiredCert(fc.Auth.DisconnectExpiredCert.Value) } - if fc.Auth.Assist != nil && fc.Auth.Assist.OpenAI != nil { - keyPath := fc.Auth.Assist.OpenAI.APITokenPath - key, err := os.ReadFile(keyPath) - if err != nil { - return trace.Errorf("failed to read OpenAI API key file: %w", err) - } - cfg.Auth.AssistAPIKey = strings.TrimSpace(string(key)) - - if fc.Auth.Assist.CommandExecutionWorkers < 0 { - return trace.BadParameter("command_execution_workers must not be negative") - } - } - // Set cluster audit configuration from file configuration. auditConfigSpec, err := services.ClusterAuditConfigSpecFromObject(fc.Storage.Params) if err != nil { @@ -987,23 +974,18 @@ func applyAuthConfig(fc *FileConfig, cfg *servicecfg.Config) error { // Only override networking configuration if some of its fields are // specified in file configuration. if fc.Auth.hasCustomNetworkingConfig() { - var assistCommandExecutionWorkers int32 - if fc.Auth.Assist != nil { - assistCommandExecutionWorkers = fc.Auth.Assist.CommandExecutionWorkers - } cfg.Auth.NetworkingConfig, err = types.NewClusterNetworkingConfigFromConfigFile(types.ClusterNetworkingConfigSpecV2{ - ClientIdleTimeout: fc.Auth.ClientIdleTimeout, - ClientIdleTimeoutMessage: fc.Auth.ClientIdleTimeoutMessage, - WebIdleTimeout: fc.Auth.WebIdleTimeout, - KeepAliveInterval: fc.Auth.KeepAliveInterval, - KeepAliveCountMax: fc.Auth.KeepAliveCountMax, - SessionControlTimeout: fc.Auth.SessionControlTimeout, - ProxyListenerMode: fc.Auth.ProxyListenerMode, - RoutingStrategy: fc.Auth.RoutingStrategy, - TunnelStrategy: fc.Auth.TunnelStrategy, - ProxyPingInterval: fc.Auth.ProxyPingInterval, - AssistCommandExecutionWorkers: assistCommandExecutionWorkers, - CaseInsensitiveRouting: fc.Auth.CaseInsensitiveRouting, + ClientIdleTimeout: fc.Auth.ClientIdleTimeout, + ClientIdleTimeoutMessage: fc.Auth.ClientIdleTimeoutMessage, + WebIdleTimeout: fc.Auth.WebIdleTimeout, + KeepAliveInterval: fc.Auth.KeepAliveInterval, + KeepAliveCountMax: fc.Auth.KeepAliveCountMax, + SessionControlTimeout: fc.Auth.SessionControlTimeout, + ProxyListenerMode: fc.Auth.ProxyListenerMode, + RoutingStrategy: fc.Auth.RoutingStrategy, + TunnelStrategy: fc.Auth.TunnelStrategy, + ProxyPingInterval: fc.Auth.ProxyPingInterval, + CaseInsensitiveRouting: fc.Auth.CaseInsensitiveRouting, }) if err != nil { return trace.Wrap(err) @@ -1238,17 +1220,6 @@ func applyProxyConfig(fc *FileConfig, cfg *servicecfg.Config) error { } } - if fc.Proxy.Assist != nil && fc.Proxy.Assist.OpenAI != nil { - keyPath := fc.Proxy.Assist.OpenAI.APITokenPath - key, err := os.ReadFile(keyPath) - if err != nil { - return trace.BadParameter("failed to read OpenAI API key file at path %s: %v", - keyPath, trace.ConvertSystemError(err)) - } else { - cfg.Proxy.AssistAPIKey = strings.TrimSpace(string(key)) - } - } - if fc.Proxy.MySQLServerVersion != "" { cfg.Proxy.MySQLServerVersion = fc.Proxy.MySQLServerVersion } diff --git a/lib/config/configuration_test.go b/lib/config/configuration_test.go index d3448ba4107..4159f78e4c3 100644 --- a/lib/config/configuration_test.go +++ b/lib/config/configuration_test.go @@ -4124,65 +4124,6 @@ func TestApplyOktaConfig(t *testing.T) { } } -func TestAssistKey(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - input string - expectKey string - expectError bool - }{ - { - desc: "api token is set", - input: ` -teleport: -proxy_service: - assist: - openai: - api_token_path: testdata/test-api-key -`, - expectKey: "123-abc-zzz", - }, - { - desc: "api token file does not exist", - input: ` -teleport: -proxy_service: - assist: - openai: - api_token_path: testdata/non-existent-file -`, - expectError: true, - }, - { - desc: "missing api token doesn't error", - input: ` -teleport: -proxy_service: - assist: - openai: -`, - expectKey: "", - }, - } { - t.Run(tc.desc, func(t *testing.T) { - conf, err := ReadConfig(strings.NewReader(tc.input)) - require.NoError(t, err) - - cfg := servicecfg.MakeDefaultConfig() - err = ApplyFileConfig(conf, cfg) - - if tc.expectError { - require.Error(t, err) - return - } - - require.Equal(t, tc.expectKey, cfg.Proxy.AssistAPIKey) - }) - } -} - func TestApplyKubeConfig(t *testing.T) { t.Parallel() diff --git a/lib/config/fileconf.go b/lib/config/fileconf.go index 4f1d0e9cbee..af4fcd6e2c4 100644 --- a/lib/config/fileconf.go +++ b/lib/config/fileconf.go @@ -806,9 +806,6 @@ type Auth struct { // This is currently Cloud-specific. HostedPlugins HostedPlugins `yaml:"hosted_plugins,omitempty"` - // Assist is a set of options related to the Teleport Assist feature. - Assist *AuthAssistOptions `yaml:"assist,omitempty"` - // AccessMonitoring is a set of options related to the Access Monitoring feature. AccessMonitoring *servicecfg.AccessMonitoringOptions `yaml:"access_monitoring,omitempty"` } @@ -851,8 +848,7 @@ func (a *Auth) hasCustomNetworkingConfig() bool { a.ProxyListenerMode != empty.ProxyListenerMode || a.RoutingStrategy != empty.RoutingStrategy || a.TunnelStrategy != empty.TunnelStrategy || - a.ProxyPingInterval != empty.ProxyPingInterval || - (a.Assist != nil && a.Assist.CommandExecutionWorkers != 0) + a.ProxyPingInterval != empty.ProxyPingInterval } // hasCustomSessionRecording returns true if any of the session recording @@ -1280,31 +1276,6 @@ func (h *HardwareKeySerialNumberValidation) Parse() (*types.HardwareKeySerialNum }, nil } -// AssistOptions is a set of options common to both Auth and Proxy related to the Teleport Assist feature. -type AssistOptions struct { - // OpenAI is a set of options related to the OpenAI assist backend. - OpenAI *OpenAIOptions `yaml:"openai,omitempty"` -} - -// ProxyAssistOptions is a set of proxy service options related to the Assist feature -type ProxyAssistOptions struct { - AssistOptions `yaml:",inline"` -} - -// AuthAssistOptions is a set of auth service options related to the Assist feature -type AuthAssistOptions struct { - AssistOptions `yaml:",inline"` - // CommandExecutionWorkers determines the number of workers that will - // execute arbitrary remote commands on servers (e.g. through Assist) in parallel - CommandExecutionWorkers int32 `yaml:"command_execution_workers,omitempty"` -} - -// OpenAIOptions stores options related to the OpenAI assist backend. -type OpenAIOptions struct { - // APITokenPath is the path to a file with OpenAI API key. - APITokenPath string `yaml:"api_token_path,omitempty"` -} - // HostedPlugins defines 'auth_service/plugins' Enterprise extension type HostedPlugins struct { Enabled bool `yaml:"enabled"` @@ -2124,9 +2095,6 @@ type Proxy struct { // UI provides config options for the web UI UI *UIConfig `yaml:"ui,omitempty"` - // Assist is a set of options related to the Teleport Assist feature. - Assist *ProxyAssistOptions `yaml:"assist,omitempty"` - // TrustXForwardedFor enables the service to take client source IPs from // the "X-Forwarded-For" headers for web APIs received from layer 7 load // balancers or reverse proxies. diff --git a/lib/modules/modules.go b/lib/modules/modules.go index 2f3a9fa3abb..affaa6a4f5f 100644 --- a/lib/modules/modules.go +++ b/lib/modules/modules.go @@ -83,8 +83,6 @@ type Features struct { AutomaticUpgrades bool // IsUsageBasedBilling enables some usage-based billing features IsUsageBasedBilling bool - // Assist enables Assistant feature - Assist bool // DeviceTrust holds its namesake feature settings. DeviceTrust DeviceTrustFeature // FeatureHiding enables hiding features from being discoverable for users who don't have the necessary permissions. @@ -191,7 +189,6 @@ func (f Features) ToProto() *proto.Features { Plugins: f.Plugins, AutomaticUpgrades: f.AutomaticUpgrades, IsUsageBased: f.IsUsageBasedBilling, - Assist: f.Assist, FeatureHiding: f.FeatureHiding, CustomTheme: f.CustomTheme, AccessGraph: f.AccessGraph, @@ -413,7 +410,6 @@ func (p *defaultModules) Features() Features { App: true, Desktop: true, AutomaticUpgrades: p.automaticUpgrades, - Assist: true, JoinActiveSessions: true, SupportType: proto.SupportType_SUPPORT_TYPE_FREE, } diff --git a/lib/service/proxy_settings.go b/lib/service/proxy_settings.go index bb8f75505c4..a87bad7ddce 100644 --- a/lib/service/proxy_settings.go +++ b/lib/service/proxy_settings.go @@ -47,12 +47,6 @@ type proxySettings struct { accessPoint networkConfigGetter } -// GetOpenAIAPIKey returns the OpenAI API key. -// TODO(jakule): Remove once plugin support is added to OSS. -func (p *proxySettings) GetOpenAIAPIKey() string { - return p.cfg.Proxy.AssistAPIKey -} - // GetProxySettings allows returns current proxy configuration. func (p *proxySettings) GetProxySettings(ctx context.Context) (*webclient.ProxySettings, error) { resp, err := p.accessPoint.GetClusterNetworkingConfig(ctx) @@ -75,7 +69,6 @@ func (p *proxySettings) GetProxySettings(ctx context.Context) (*webclient.ProxyS func (p *proxySettings) buildProxySettings(proxyListenerMode types.ProxyListenerMode) *webclient.ProxySettings { proxySettings := webclient.ProxySettings{ TLSRoutingEnabled: proxyListenerMode == types.ProxyListenerMode_Multiplex, - AssistEnabled: p.cfg.Proxy.AssistAPIKey != "", Kube: webclient.KubeProxySettings{ Enabled: p.cfg.Proxy.Kube.Enabled, }, diff --git a/lib/service/proxy_settings_test.go b/lib/service/proxy_settings_test.go deleted file mode 100644 index a1f376776f1..00000000000 --- a/lib/service/proxy_settings_test.go +++ /dev/null @@ -1,111 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package service - -import ( - "context" - "testing" - - "github.com/stretchr/testify/require" - - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/lib/defaults" - "github.com/gravitational/teleport/lib/service/servicecfg" - "github.com/gravitational/teleport/lib/utils" -) - -type accessPointMock struct{} - -// GetClusterNetworkingConfig returns a cluster config. -func (a *accessPointMock) GetClusterNetworkingConfig(_ context.Context) (types.ClusterNetworkingConfig, error) { - return &types.ClusterNetworkingConfigV2{ - Spec: types.ClusterNetworkingConfigSpecV2{ - RoutingStrategy: types.RoutingStrategy_MOST_RECENT, - }, - }, nil - -} - -func Test_proxySettings_GetProxySettings(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - cfgFn func() *servicecfg.Config - assertFn require.BoolAssertionFunc - }{ - { - name: "AssistEnabled is true when proxy API key is set", - cfgFn: func() *servicecfg.Config { - cfg := servicecfg.MakeDefaultConfig() - cfg.Proxy.AssistAPIKey = "test-api-key" - return cfg - }, - assertFn: require.True, - }, - { - name: "AssistEnabled is false when proxy API key is not set", - cfgFn: func() *servicecfg.Config { - cfg := servicecfg.MakeDefaultConfig() - cfg.Proxy.AssistAPIKey = "" - return cfg - }, - assertFn: require.False, - }, - { - name: "AssistEnabled is true when proxy API key is set - v2 config", - cfgFn: func() *servicecfg.Config { - cfg := servicecfg.MakeDefaultConfig() - cfg.Version = defaults.TeleportConfigVersionV2 - cfg.Proxy.AssistAPIKey = "test-api-key" - return cfg - }, - assertFn: require.True, - }, - { - name: "AssistEnabled is false when proxy API key is not set - v2 config", - cfgFn: func() *servicecfg.Config { - cfg := servicecfg.MakeDefaultConfig() - cfg.Version = defaults.TeleportConfigVersionV2 - cfg.Proxy.AssistAPIKey = "" - return cfg - }, - assertFn: require.False, - }, - } - - for _, tt := range tests { - tt := tt - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - - p := &proxySettings{ - cfg: tt.cfgFn(), - proxySSHAddr: utils.NetAddr{AddrNetwork: "tcp", Addr: "0.0.0.0:3080"}, - accessPoint: &accessPointMock{}, - } - - proxySettings, err := p.GetProxySettings(context.Background()) - require.NoError(t, err) - require.NotNil(t, proxySettings) - - tt.assertFn(t, proxySettings.AssistEnabled) - }) - } -} diff --git a/lib/service/service.go b/lib/service/service.go index 0ffd1b67253..db6c28add49 100644 --- a/lib/service/service.go +++ b/lib/service/service.go @@ -79,11 +79,8 @@ import ( apiutils "github.com/gravitational/teleport/api/utils" "github.com/gravitational/teleport/api/utils/aws" "github.com/gravitational/teleport/api/utils/grpc/interceptors" - "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib" "github.com/gravitational/teleport/lib/agentless" - "github.com/gravitational/teleport/lib/ai" - "github.com/gravitational/teleport/lib/ai/embedding" "github.com/gravitational/teleport/lib/auditd" "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/auth/accesspoint" @@ -287,15 +284,6 @@ const ( TeleportOKEvent = "TeleportOKEvent" ) -const ( - // embeddingInitialDelay is the time to wait before the first embedding - // routine is started. - embeddingInitialDelay = 10 * time.Second - // embeddingPeriod is the time between two embedding routines. - // A seventh jitter is applied on the period. - embeddingPeriod = 20 * time.Minute -) - // Connector has all resources process needs to connect to other parts of the // cluster: client and identity. type Connector struct { @@ -1870,19 +1858,6 @@ func (process *TeleportProcess) initAuthService() error { traceClt = clt } - var embedderClient embedding.Embedder - if cfg.Auth.AssistAPIKey != "" { - // cfg.Testing.OpenAIConfig is set in tests to change the OpenAI API endpoint - // Like for proxy, if a custom OpenAIConfig is passed, the token from - // cfg.Auth.AssistAPIKey is ignored and the one from the config is used. - if cfg.Testing.OpenAIConfig != nil { - embedderClient = ai.NewClientFromConfig(*cfg.Testing.OpenAIConfig) - } else { - embedderClient = ai.NewClient(cfg.Auth.AssistAPIKey) - } - } - - embeddingsRetriever := ai.NewSimpleRetriever() cn, err := services.NewClusterNameWithRandomID(types.ClusterNameSpecV2{ ClusterName: clusterName, }) @@ -1961,8 +1936,6 @@ func (process *TeleportProcess) initAuthService() error { AccessMonitoringEnabled: cfg.Auth.IsAccessMonitoringEnabled(), Clock: cfg.Clock, HTTPClientForAWSSTS: cfg.Auth.HTTPClientForAWSSTS, - EmbeddingRetriever: embeddingsRetriever, - EmbeddingClient: embedderClient, Tracer: process.TracingProvider.Tracer(teleport.ComponentAuth), CloudClients: cloudClients, }, func(as *auth.Server) error { @@ -2052,40 +2025,6 @@ func (process *TeleportProcess) initAuthService() error { authServer.SetGlobalNotificationCache(globalNotificationCache) - if embedderClient != nil { - logger.DebugContext(process.ExitContext(), "Starting embedding watcher") - embeddingProcessor := ai.NewEmbeddingProcessor(&ai.EmbeddingProcessorConfig{ - AIClient: embedderClient, - EmbeddingsRetriever: embeddingsRetriever, - EmbeddingSrv: authServer, - NodeSrv: authServer.UnifiedResourceCache, - Log: process.log.WithField(teleport.ComponentKey, teleport.Component(teleport.ComponentAuth, process.id)), - Jitter: retryutils.NewFullJitter(), - }) - - process.RegisterFunc("ai.embedding-processor", func() error { - // We check the Assist feature flag here rather than on creation of TeleportProcess, - // as when running Enterprise and the feature source is Cloud, - // features may be loaded at two different times: - // 1. When Cloud is reachable, features will be fetched from Cloud - // before constructing TeleportProcess - // 2. When Cloud is not reachable, we will attempt to load cached features - // from the Teleport backend. - // In the second case, we don't know the final value of Features().Assist - // when constructing the process. - // Services in the supervisor will only start after either 1 or 2 has succeeded, - // so we can make the decision here. - // - // Ref: e/tool/teleport/process/process.go - if !modules.GetModules().Features().Assist { - logger.DebugContext(process.ExitContext(), "Skipping start of embedding processor: Assist feature not enabled for license") - return nil - } - logger.DebugContext(process.ExitContext(), "Starting embedding processor") - return embeddingProcessor.Run(process.GracefulExitContext(), embeddingInitialDelay, embeddingPeriod) - }) - } - headlessAuthenticationWatcher, err := local.NewHeadlessAuthenticationWatcher(process.ExitContext(), local.HeadlessAuthenticationWatcherConfig{ Backend: b, }) @@ -4399,7 +4338,6 @@ func (process *TeleportProcess) initProxyEndpoint(conn *Connector) error { return ctx, trace.Wrap(err) }), PROXYSigner: proxySigner, - OpenAIConfig: cfg.Testing.OpenAIConfig, NodeWatcher: nodeWatcher, AccessGraphAddr: accessGraphAddr, TracerProvider: process.TracingProvider, diff --git a/lib/service/servicecfg/auth.go b/lib/service/servicecfg/auth.go index 93a4d2eb57f..40277cd81e6 100644 --- a/lib/service/servicecfg/auth.go +++ b/lib/service/servicecfg/auth.go @@ -113,10 +113,6 @@ type AuthConfig struct { // STS requests. Used in test. HTTPClientForAWSSTS utils.HTTPDoClient - // AssistAPIKey is the OpenAI API key. - // TODO: This key will be moved to a plugin once support for plugins is implemented. - AssistAPIKey string - // AccessMonitoring configures access monitoring. AccessMonitoring *AccessMonitoringOptions } diff --git a/lib/service/servicecfg/config.go b/lib/service/servicecfg/config.go index f339244f2f1..456321cdf35 100644 --- a/lib/service/servicecfg/config.go +++ b/lib/service/servicecfg/config.go @@ -33,7 +33,6 @@ import ( "github.com/ghodss/yaml" "github.com/gravitational/trace" "github.com/jonboulle/clockwork" - "github.com/sashabaranov/go-openai" "github.com/sirupsen/logrus" "golang.org/x/crypto/ssh" @@ -304,15 +303,6 @@ type ConfigTesting struct { // require PROXY header if 'proxyProtocolMode: true' even from self connections. Used in tests as all connections are self // connections there. KubeMultiplexerIgnoreSelfConnections bool - - // OpenAIConfig contains the optional OpenAI client configuration used by - // auth and proxy. When it's not set (the default, we don't offer a way to - // set it when executing the regular Teleport binary) we use the default - // configuration with auth tokens passed from Auth.AssistAPIKey or - // Proxy.AssistAPIKey. We set this only when testing to avoid calls to reach - // the real OpenAI API. - // Note: When set, this overrides Auth and Proxy's AssistAPIKey settings. - OpenAIConfig *openai.ClientConfig } // AccessGraphConfig represents TAG server config diff --git a/lib/service/servicecfg/proxy.go b/lib/service/servicecfg/proxy.go index 70be8d9aa8f..c07ce5d47b0 100644 --- a/lib/service/servicecfg/proxy.go +++ b/lib/service/servicecfg/proxy.go @@ -139,10 +139,6 @@ type ProxyConfig struct { // UI provides config options for the web UI UI webclient.UIConfig - // AssistAPIKey is the OpenAI API key. - // TODO: This key will be moved to a plugin once support for plugins is implemented. - AssistAPIKey string - // TrustXForwardedFor enables the service to take client source IPs from // the "X-Forwarded-For" headers for web APIs recevied from layer 7 load // balancers or reverse proxies. diff --git a/lib/services/assist.go b/lib/services/assist.go deleted file mode 100644 index fa60cf26c5f..00000000000 --- a/lib/services/assist.go +++ /dev/null @@ -1,48 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package services - -import ( - "context" - - "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" -) - -type Assistant interface { - // GetAssistantMessages returns all messages with given conversation ID. - GetAssistantMessages(ctx context.Context, req *assist.GetAssistantMessagesRequest) (*assist.GetAssistantMessagesResponse, error) - - // CreateAssistantMessage adds the message to the backend. - CreateAssistantMessage(ctx context.Context, msg *assist.CreateAssistantMessageRequest) error - - // CreateAssistantConversation creates a new conversation entry in the backend. - CreateAssistantConversation(ctx context.Context, req *assist.CreateAssistantConversationRequest) (*assist.CreateAssistantConversationResponse, error) - - // DeleteAssistantConversation deletes a conversation entry and associated messages from the backend. - DeleteAssistantConversation(ctx context.Context, req *assist.DeleteAssistantConversationRequest) error - - // GetAssistantConversations returns all conversations started by a user. - GetAssistantConversations(ctx context.Context, request *assist.GetAssistantConversationsRequest) (*assist.GetAssistantConversationsResponse, error) - - // UpdateAssistantConversationInfo updates conversation info. - UpdateAssistantConversationInfo(ctx context.Context, msg *assist.UpdateAssistantConversationInfoRequest) error - - // IsAssistEnabled returns true if the assist is enabled or not on the auth level. - IsAssistEnabled(ctx context.Context) (*assist.IsAssistEnabledResponse, error) -} diff --git a/lib/services/embeddings.go b/lib/services/embeddings.go deleted file mode 100644 index db227832a60..00000000000 --- a/lib/services/embeddings.go +++ /dev/null @@ -1,40 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package services - -import ( - "context" - - "github.com/gravitational/teleport/api/internalutils/stream" - "github.com/gravitational/teleport/lib/ai/embedding" -) - -// Embeddings service is responsible for storing and retrieving embeddings in -// the backend. The backend acts as an embedding cache. Embeddings can be -// re-generated by an ai.Embedder. -type Embeddings interface { - // GetEmbedding looks up a single embedding by its name in the backend. - GetEmbedding(ctx context.Context, kind, resourceID string) (*embedding.Embedding, error) - // GetEmbeddings returns all embeddings for a given kind. - GetEmbeddings(ctx context.Context, kind string) stream.Stream[*embedding.Embedding] - // GetEmbeddings returns all embeddings. - GetAllEmbeddings(ctx context.Context) stream.Stream[*embedding.Embedding] - // UpsertEmbedding creates or updates a single ai.Embedding in the backend. - UpsertEmbedding(ctx context.Context, embedding *embedding.Embedding) (*embedding.Embedding, error) -} diff --git a/lib/services/local/assistant.go b/lib/services/local/assistant.go deleted file mode 100644 index 37eae0c5b97..00000000000 --- a/lib/services/local/assistant.go +++ /dev/null @@ -1,285 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package local - -import ( - "context" - "encoding/json" - "sort" - "time" - - "github.com/google/uuid" - "github.com/gravitational/trace" - "github.com/sirupsen/logrus" - "google.golang.org/protobuf/types/known/timestamppb" - - "github.com/gravitational/teleport" - "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" - "github.com/gravitational/teleport/lib/backend" -) - -// Conversation is a conversation entry in the backend. -type Conversation struct { - Title string `json:"title,omitempty"` - ConversationID string `json:"conversation_id"` - CreatedTime time.Time `json:"created_time"` -} - -// AssistService is responsible for managing assist conversations. -type AssistService struct { - backend.Backend - log logrus.FieldLogger -} - -// NewAssistService returns a new instance of AssistService. -func NewAssistService(backend backend.Backend) *AssistService { - return &AssistService{ - Backend: backend, - log: logrus.WithField(teleport.ComponentKey, "assist"), - } -} - -// CreateAssistantConversation creates a new conversation entry in the backend. -func (s *AssistService) CreateAssistantConversation(ctx context.Context, - req *assist.CreateAssistantConversationRequest, -) (*assist.CreateAssistantConversationResponse, error) { - if req.Username == "" { - return nil, trace.BadParameter("missing parameter username") - } - if req.CreatedTime == nil { - return nil, trace.BadParameter("missing parameter created time") - } - - conversationID := uuid.New().String() - payload := &Conversation{ - ConversationID: conversationID, - CreatedTime: req.GetCreatedTime().AsTime(), - } - - value, err := json.Marshal(payload) - if err != nil { - return nil, trace.Wrap(err) - } - - item := backend.Item{ - Key: backend.Key(assistantConversationPrefix, req.Username, conversationID), - Value: value, - } - - _, err = s.Create(ctx, item) - if err != nil { - return nil, trace.Wrap(err) - } - - return &assist.CreateAssistantConversationResponse{Id: conversationID}, nil -} - -// getConversation returns a conversation from the backend. -func (s *AssistService) getConversation(ctx context.Context, username, conversationID string) (*Conversation, error) { - item, err := s.Get(ctx, backend.Key(assistantConversationPrefix, username, conversationID)) - if err != nil { - return nil, trace.Wrap(err) - } - - var conversation Conversation - if err := json.Unmarshal(item.Value, &conversation); err != nil { - return nil, trace.Wrap(err) - } - - return &conversation, nil -} - -// DeleteAssistantConversation deletes a conversation from the backend. -func (s *AssistService) DeleteAssistantConversation(ctx context.Context, req *assist.DeleteAssistantConversationRequest) error { - if req.Username == "" { - return trace.BadParameter("missing parameter username") - } - if req.ConversationId == "" { - return trace.BadParameter("missing parameter conversation ID") - } - - // Delete all messages in the conversation first, so that if the delete - // fails, the conversation is still there. Client can retry the deleting. - if err := s.deleteAllMessages(ctx, req.Username, req.ConversationId); err != nil { - return trace.Wrap(err) - } - - // Delete the conversation. - if err := s.Delete(ctx, backend.Key(assistantConversationPrefix, req.Username, req.ConversationId)); err != nil { - return trace.Wrap(err) - } - - return nil -} - -// deleteAllMessages deletes all messages in a conversation. -func (s *AssistService) deleteAllMessages(ctx context.Context, username, conversationID string) error { - startKey := backend.ExactKey(assistantMessagePrefix, username, conversationID) - if err := s.DeleteRange(ctx, startKey, backend.RangeEnd(startKey)); err != nil { - return trace.Wrap(err) - } - - return nil -} - -// UpdateAssistantConversationInfo updates the conversation title. -func (s *AssistService) UpdateAssistantConversationInfo(ctx context.Context, request *assist.UpdateAssistantConversationInfoRequest) error { - if request.ConversationId == "" { - return trace.BadParameter("missing conversation ID") - } - if request.Username == "" { - return trace.BadParameter("missing username") - } - if request.Title == "" { - return trace.BadParameter("missing title") - } - - msg, err := s.getConversation(ctx, request.Username, request.GetConversationId()) - if err != nil { - return trace.Wrap(err) - } - - // Only update the title, leave the rest of the fields intact. - msg.Title = request.Title - - payload, err := json.Marshal(msg) - if err != nil { - return trace.Wrap(err) - } - - item := backend.Item{ - Key: backend.Key(assistantConversationPrefix, request.GetUsername(), request.GetConversationId()), - Value: payload, - } - - if _, err = s.Update(ctx, item); err != nil { - return trace.Wrap(err) - } - - return nil -} - -// GetAssistantConversations returns all conversations started by a user. -func (s *AssistService) GetAssistantConversations(ctx context.Context, req *assist.GetAssistantConversationsRequest) (*assist.GetAssistantConversationsResponse, error) { - if req.Username == "" { - return nil, trace.BadParameter("missing username") - } - startKey := backend.ExactKey(assistantConversationPrefix, req.Username) - result, err := s.GetRange(ctx, startKey, backend.RangeEnd(startKey), backend.NoLimit) - if err != nil { - return nil, trace.Wrap(err) - } - - conversationsIDs := make([]*assist.ConversationInfo, 0, len(result.Items)) - for _, item := range result.Items { - payload := &Conversation{} - if err := json.Unmarshal(item.Value, payload); err != nil { - return nil, trace.Wrap(err) - } - - conversationsIDs = append(conversationsIDs, &assist.ConversationInfo{ - Id: payload.ConversationID, - Title: payload.Title, - CreatedTime: timestamppb.New(payload.CreatedTime), - }) - } - - sort.Slice(conversationsIDs, func(i, j int) bool { - return conversationsIDs[i].CreatedTime.AsTime().Before(conversationsIDs[j].GetCreatedTime().AsTime()) - }) - - return &assist.GetAssistantConversationsResponse{ - Conversations: conversationsIDs, - }, nil -} - -// GetAssistantMessages returns all messages with given conversation ID. -func (s *AssistService) GetAssistantMessages(ctx context.Context, req *assist.GetAssistantMessagesRequest) (*assist.GetAssistantMessagesResponse, error) { - if req.Username == "" { - return nil, trace.BadParameter("missing username") - } - - if req.ConversationId == "" { - return nil, trace.BadParameter("missing conversation ID") - } - - startKey := backend.ExactKey(assistantMessagePrefix, req.Username, req.ConversationId) - result, err := s.GetRange(ctx, startKey, backend.RangeEnd(startKey), backend.NoLimit) - if err != nil { - return nil, trace.Wrap(err) - } - - out := make([]*assist.AssistantMessage, len(result.Items)) - for i, item := range result.Items { - var a assist.AssistantMessage - if err := json.Unmarshal(item.Value, &a); err != nil { - return nil, trace.Wrap(err) - } - out[i] = &a - } - - sort.Slice(out, func(i, j int) bool { - // Sort by created time. - return out[i].CreatedTime.AsTime().Before(out[j].GetCreatedTime().AsTime()) - }) - - return &assist.GetAssistantMessagesResponse{ - Messages: out, - }, nil -} - -// CreateAssistantMessage adds the message to the backend. -func (s *AssistService) CreateAssistantMessage(ctx context.Context, req *assist.CreateAssistantMessageRequest) error { - if req.Username == "" { - return trace.BadParameter("missing username") - } - if req.ConversationId == "" { - return trace.BadParameter("missing conversation ID") - } - - // Check if the conversation exists. - conversationKey := backend.Key(assistantConversationPrefix, req.Username, req.ConversationId) - if _, err := s.Get(ctx, conversationKey); err != nil { - if trace.IsNotFound(err) { - return trace.NotFound("conversation %q not found", req.ConversationId) - } - return trace.Wrap(err) - } - - msg := req.GetMessage() - value, err := json.Marshal(msg) - if err != nil { - return trace.Wrap(err) - } - - messageID := uuid.New().String() - - item := backend.Item{ - Key: backend.Key(assistantMessagePrefix, req.Username, req.ConversationId, messageID), - Value: value, - } - - _, err = s.Create(ctx, item) - return trace.Wrap(err) -} - -// IsAssistEnabled returns true if the assist is enabled or not on the auth level. -func (s *AssistService) IsAssistEnabled(ctx context.Context) (*assist.IsAssistEnabledResponse, error) { - return &assist.IsAssistEnabledResponse{Enabled: s.Backend.GetName() != "etcd"}, nil -} diff --git a/lib/services/local/assistant_test.go b/lib/services/local/assistant_test.go deleted file mode 100644 index 1fe8d06f0b4..00000000000 --- a/lib/services/local/assistant_test.go +++ /dev/null @@ -1,176 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package local_test - -import ( - "context" - "testing" - "time" - - "github.com/google/uuid" - "github.com/jonboulle/clockwork" - "github.com/stretchr/testify/require" - "google.golang.org/protobuf/types/known/timestamppb" - - "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" - "github.com/gravitational/teleport/lib/backend/memory" - "github.com/gravitational/teleport/lib/services/local" -) - -func newAssistService(t *testing.T) *local.AssistService { - t.Helper() - backend, err := memory.New(memory.Config{ - Context: context.Background(), - Clock: clockwork.NewFakeClock(), - }) - require.NoError(t, err) - return local.NewAssistService(backend) -} - -// TestAssistantCRUD tests the assistant CRUD operations. -func TestAssistantCRUD(t *testing.T) { - t.Parallel() - - identity := newAssistService(t) - ctx := context.Background() - - const username = "foo" - var conversationID string - - t.Run("create conversation", func(t *testing.T) { - req := &assist.CreateAssistantConversationRequest{ - Username: username, - CreatedTime: timestamppb.New(time.Now()), - } - - conversationResp, err := identity.CreateAssistantConversation(ctx, req) - require.NoError(t, err) - require.NotEmpty(t, conversationResp.Id) - - conversationID = conversationResp.Id - }) - - t.Run("get conversation", func(t *testing.T) { - req := &assist.GetAssistantConversationsRequest{ - Username: username, - } - conversations, err := identity.GetAssistantConversations(ctx, req) - require.NoError(t, err) - require.Len(t, conversations.Conversations, 1) - require.Equal(t, conversationID, conversations.Conversations[0].Id) - }) - - t.Run("create message", func(t *testing.T) { - msg := &assist.CreateAssistantMessageRequest{ - Username: username, - ConversationId: conversationID, - Message: &assist.AssistantMessage{ - CreatedTime: timestamppb.New(time.Now()), - Payload: "foo", - Type: "USER_MSG", - }, - } - err := identity.CreateAssistantMessage(ctx, msg) - require.NoError(t, err) - }) - - t.Run("get messages", func(t *testing.T) { - req := &assist.GetAssistantMessagesRequest{ - Username: username, - ConversationId: conversationID, - } - messages, err := identity.GetAssistantMessages(ctx, req) - require.NoError(t, err) - require.Len(t, messages.Messages, 1) - require.Equal(t, "foo", messages.Messages[0].Payload) - }) - - t.Run("set conversation title", func(t *testing.T) { - titleReq := &assist.UpdateAssistantConversationInfoRequest{ - Title: "bar", - Username: username, - ConversationId: conversationID, - } - title := "bar" - err := identity.UpdateAssistantConversationInfo(ctx, titleReq) - require.NoError(t, err) - - req := &assist.GetAssistantConversationsRequest{ - Username: username, - } - - conversations, err := identity.GetAssistantConversations(ctx, req) - require.NoError(t, err) - require.Len(t, conversations.Conversations, 1) - require.Equal(t, title, conversations.Conversations[0].Title) - }) - - t.Run("conversations are sorted by created_time", func(t *testing.T) { - req := &assist.CreateAssistantConversationRequest{ - Username: username, - CreatedTime: timestamppb.New(time.Now().Add(time.Hour)), - } - - conversationResp, err := identity.CreateAssistantConversation(ctx, req) - require.NoError(t, err) - require.NotEmpty(t, conversationResp.Id) - - reqConversations := &assist.GetAssistantConversationsRequest{ - Username: username, - } - - conversations, err := identity.GetAssistantConversations(ctx, reqConversations) - require.NoError(t, err) - require.Len(t, conversations.Conversations, 2) - require.Equal(t, conversationID, conversations.Conversations[0].Id) - require.Equal(t, conversationResp.Id, conversations.Conversations[1].Id) - }) - - t.Run("refuse to add messages if conversion does not exist", func(t *testing.T) { - msg := &assist.CreateAssistantMessageRequest{ - Username: username, - ConversationId: uuid.New().String(), - Message: &assist.AssistantMessage{ - CreatedTime: timestamppb.New(time.Now()), - Payload: "foo", - Type: "USER_MSG", - }, - } - err := identity.CreateAssistantMessage(ctx, msg) - require.Error(t, err) - }) - - t.Run("delete conversation", func(t *testing.T) { - req := &assist.DeleteAssistantConversationRequest{ - Username: username, - ConversationId: conversationID, - } - err := identity.DeleteAssistantConversation(ctx, req) - require.NoError(t, err) - - reqConversations := &assist.GetAssistantConversationsRequest{ - Username: username, - } - - conversations, err := identity.GetAssistantConversations(ctx, reqConversations) - require.NoError(t, err) - require.Len(t, conversations.Conversations, 1) - require.NotEqual(t, conversationID, conversations.Conversations[0].Id, "conversation was not deleted") - }) -} diff --git a/lib/services/local/embeddings.go b/lib/services/local/embeddings.go deleted file mode 100644 index 1afdfb35a18..00000000000 --- a/lib/services/local/embeddings.go +++ /dev/null @@ -1,117 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package local - -import ( - "context" - "time" - - "github.com/gravitational/trace" - "github.com/jonboulle/clockwork" - "github.com/sirupsen/logrus" - - "github.com/gravitational/teleport" - "github.com/gravitational/teleport/api/internalutils/stream" - "github.com/gravitational/teleport/api/utils/retryutils" - "github.com/gravitational/teleport/lib/ai" - "github.com/gravitational/teleport/lib/ai/embedding" - "github.com/gravitational/teleport/lib/backend" -) - -// EmbeddingsService implements the services.Embeddings interface. -type EmbeddingsService struct { - log *logrus.Entry - jitter retryutils.Jitter - backend.Backend - clock clockwork.Clock -} - -const ( - embeddingsPrefix = "embeddings" - embeddingExpiry = 30 * 24 * time.Hour // 30 days -) - -// GetEmbedding looks up a single embedding by its name in the backend. -func (e EmbeddingsService) GetEmbedding(ctx context.Context, kind, resourceID string) (*embedding.Embedding, error) { - result, err := e.Get(ctx, backend.Key(embeddingsPrefix, kind, resourceID)) - if err != nil { - return nil, trace.Wrap(err) - } - return ai.UnmarshalEmbedding(result.Value) -} - -// GetEmbeddings returns a stream of all embeddings -func (e EmbeddingsService) GetAllEmbeddings(ctx context.Context) stream.Stream[*embedding.Embedding] { - startKey := backend.ExactKey(embeddingsPrefix) - items := backend.StreamRange(ctx, e, startKey, backend.RangeEnd(startKey), 50) - return stream.FilterMap(items, func(item backend.Item) (*embedding.Embedding, bool) { - embedding, err := ai.UnmarshalEmbedding(item.Value) - if err != nil { - e.log.Warnf("Skipping embedding at %s, failed to unmarshal: %v", item.Key, err) - return nil, false - } - return embedding, true - }) -} - -// GetEmbeddings returns a stream of embeddings for a given kind. -func (e EmbeddingsService) GetEmbeddings(ctx context.Context, kind string) stream.Stream[*embedding.Embedding] { - startKey := backend.ExactKey(embeddingsPrefix, kind) - items := backend.StreamRange(ctx, e, startKey, backend.RangeEnd(startKey), 50) - return stream.FilterMap(items, func(item backend.Item) (*embedding.Embedding, bool) { - embedding, err := ai.UnmarshalEmbedding(item.Value) - if err != nil { - e.log.Warnf("Skipping embedding at %s, failed to unmarshal: %v", item.Key, err) - return nil, false - } - return embedding, true - }) -} - -// UpsertEmbedding creates or update a single ai.Embedding in the backend. -func (e EmbeddingsService) UpsertEmbedding(ctx context.Context, embedding *embedding.Embedding) (*embedding.Embedding, error) { - value, err := ai.MarshalEmbedding(embedding) - if err != nil { - return nil, trace.Wrap(err) - } - _, err = e.Put(ctx, backend.Item{ - Key: embeddingItemKey(embedding), - Value: value, - Expires: e.clock.Now().Add(embeddingExpiry), - }) - if err != nil { - return nil, trace.Wrap(err) - } - return embedding, nil -} - -// NewEmbeddingsService is a constructor for the EmbeddingsService. -func NewEmbeddingsService(b backend.Backend) *EmbeddingsService { - return &EmbeddingsService{ - log: logrus.WithFields(logrus.Fields{teleport.ComponentKey: "Embeddings"}), - jitter: retryutils.NewFullJitter(), - Backend: b, - clock: clockwork.NewRealClock(), - } -} - -// embeddingItemKey builds the backend item key for a given ai.Embedding. -func embeddingItemKey(embedding *embedding.Embedding) []byte { - return backend.Key(embeddingsPrefix, embedding.GetName()) -} diff --git a/lib/services/local/embeddings_test.go b/lib/services/local/embeddings_test.go deleted file mode 100644 index 0fd89579e3f..00000000000 --- a/lib/services/local/embeddings_test.go +++ /dev/null @@ -1,237 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package local - -import ( - "context" - "crypto/sha256" - "sort" - "testing" - - "github.com/gravitational/trace" - "github.com/jonboulle/clockwork" - "github.com/stretchr/testify/require" - - "github.com/gravitational/teleport/api/internalutils/stream" - "github.com/gravitational/teleport/api/types" - embeddinglib "github.com/gravitational/teleport/lib/ai/embedding" - "github.com/gravitational/teleport/lib/backend/memory" -) - -var ( - embedding1 = embeddinglib.NewEmbedding(types.KindNode, "foo", embeddinglib.Vector64{0, 0}, sha256.Sum256([]byte("test1"))) - embedding2 = embeddinglib.NewEmbedding(types.KindNode, "bar", embeddinglib.Vector64{1, 1, 1}, sha256.Sum256([]byte("test2"))) - embedding3 = embeddinglib.NewEmbedding(types.KindDatabase, "bar", embeddinglib.Vector64{2}, sha256.Sum256([]byte("test3"))) -) - -func errorIsNotFound(t require.TestingT, err error, msgAndArgs ...interface{}) { - require.True(t, trace.IsNotFound(err), msgAndArgs...) -} - -func TestGetEmbedding(t *testing.T) { - t.Parallel() - // Test setup: create the backend, the service, and load all fixtures - ctx := context.Background() - - fixtures := []*embeddinglib.Embedding{embedding1, embedding2, embedding3} - - backend, err := memory.New(memory.Config{ - Context: ctx, - Clock: clockwork.NewFakeClock(), - }) - require.NoError(t, err) - - service := NewEmbeddingsService(backend) - - for _, fixture := range fixtures { - _, err := service.UpsertEmbedding(ctx, fixture) - require.NoError(t, err) - } - - // Test execution - tests := []struct { - name string - kind string - id string - assertErr require.ErrorAssertionFunc - expected *embeddinglib.Embedding - }{ - { - name: "Simple get", - kind: types.KindNode, - id: "foo", - assertErr: require.NoError, - expected: embedding1, - }, - { - name: "Kind conflict", - kind: types.KindDatabase, - id: "bar", - assertErr: require.NoError, - expected: embedding3, - }, - { - name: "Non-existing", - kind: types.KindDatabase, - id: "foo", - assertErr: errorIsNotFound, - expected: nil, - }, - } - - for _, tc := range tests { - tc := tc - t.Run(tc.name, func(t *testing.T) { - t.Parallel() - embedding, err := service.GetEmbedding(ctx, tc.kind, tc.id) - tc.assertErr(t, err) - requireEmbeddingsEqual(t, tc.expected, embedding) - }) - } -} - -func TestGetEmbeddings(t *testing.T) { - t.Parallel() - // Test setup: create the backend, the service, and load all fixtures - ctx := context.Background() - - fixtures := []*embeddinglib.Embedding{embedding1, embedding2, embedding3} - - backend, err := memory.New(memory.Config{ - Context: ctx, - Clock: clockwork.NewFakeClock(), - }) - require.NoError(t, err) - - service := NewEmbeddingsService(backend) - - for _, fixture := range fixtures { - _, err := service.UpsertEmbedding(ctx, fixture) - require.NoError(t, err) - } - - // Test execution - tests := []struct { - name string - kind string - assertErr require.ErrorAssertionFunc - expected sortableEmbeddings - }{ - { - name: "Get multiple embeddings", - kind: types.KindNode, - assertErr: require.NoError, - expected: sortableEmbeddings{embedding1, embedding2}, - }, - { - name: "Get single embedding", - kind: types.KindDatabase, - assertErr: require.NoError, - expected: sortableEmbeddings{embedding3}, - }, - { - name: "Get no embeddings", - kind: types.KindApp, - assertErr: require.NoError, - expected: nil, - }, - } - for _, tc := range tests { - tc := tc - t.Run(tc.name, func(t *testing.T) { - t.Parallel() - var embeddings sortableEmbeddings - var err error - embeddings, err = stream.Collect(service.GetEmbeddings(ctx, tc.kind)) - tc.assertErr(t, err) - sort.Sort(embeddings) - sort.Sort(tc.expected) - require.Equal(t, len(tc.expected), len(embeddings)) - for i, expected := range tc.expected { - requireEmbeddingsEqual(t, expected, embeddings[i]) - } - }) - } -} - -func TestUpsertEmbedding(t *testing.T) { - t.Parallel() - // Test setup: create the backend, the service - ctx := context.Background() - - backend, err := memory.New(memory.Config{ - Context: ctx, - Clock: clockwork.NewFakeClock(), - }) - require.NoError(t, err) - - service := NewEmbeddingsService(backend) - - // Test: check there's nothing in the backend first - _, err = service.GetEmbedding(ctx, types.KindNode, "foo") - errorIsNotFound(t, err) - - // Test: add an element in the backend and check if we can retrieve it - embedding := embeddinglib.NewEmbedding(types.KindNode, "foo", embeddinglib.Vector64{0, 0}, sha256.Sum256([]byte("test"))) - embedding, err = service.UpsertEmbedding(ctx, embedding) - require.NoError(t, err) - result, err := service.GetEmbedding(ctx, types.KindNode, "foo") - require.NoError(t, err) - requireEmbeddingsEqual(t, embedding, result) - - // Test: update the embedding and check we now retrieve the new version - embedding = embeddinglib.NewEmbedding(types.KindNode, "foo", embeddinglib.Vector64{1, 1, 1, 1, 1}, sha256.Sum256([]byte("test2"))) - embedding, err = service.UpsertEmbedding(ctx, embedding) - require.NoError(t, err) - result, err = service.GetEmbedding(ctx, types.KindNode, "foo") - require.NoError(t, err) - requireEmbeddingsEqual(t, embedding, result) -} - -// sortableEmbeddings is an embedding.Embedding list that can be sorted. This is used -// in tests to compare two lists and their content. -type sortableEmbeddings []*embeddinglib.Embedding - -func (s sortableEmbeddings) Len() int { - return len(s) -} - -func (s sortableEmbeddings) Less(i, j int) bool { - return s[i].GetName() < s[j].GetName() -} - -func (s sortableEmbeddings) Swap(i, j int) { - s[i], s[j] = s[j], s[i] -} - -// requireEmbeddingsEqual checks if two embeddings are equal or fails the test otherwise. -// This is required because equivalent ai.Embedding might differ depending on -// how they have been created (marshaling/unmarshalling protobuf messages set -// some internal fields that a freshly created ai.Embedding doesn't have). -func requireEmbeddingsEqual(t require.TestingT, expected, actual *embeddinglib.Embedding) { - if expected == nil { - require.Nil(t, actual) - return - } - require.NotNil(t, actual) - require.Equal(t, expected.EmbeddedId, actual.EmbeddedId) - require.Equal(t, expected.EmbeddedKind, actual.EmbeddedKind) - require.Equal(t, expected.EmbeddedHash, actual.EmbeddedHash) - require.Equal(t, expected.Vector, actual.Vector) -} diff --git a/lib/services/local/userpreferences.go b/lib/services/local/userpreferences.go index a491494865a..a5518b80a99 100644 --- a/lib/services/local/userpreferences.go +++ b/lib/services/local/userpreferences.go @@ -36,10 +36,6 @@ type UserPreferencesService struct { func DefaultUserPreferences() *userpreferencesv1.UserPreferences { return &userpreferencesv1.UserPreferences{ - Assist: &userpreferencesv1.AssistUserPreferences{ - PreferredLogins: []string{}, - ViewMode: userpreferencesv1.AssistViewMode_ASSIST_VIEW_MODE_DOCKED, - }, Theme: userpreferencesv1.Theme_THEME_UNSPECIFIED, UnifiedResourcePreferences: &userpreferencesv1.UnifiedResourcePreferences{ DefaultTab: userpreferencesv1.DefaultTab_DEFAULT_TAB_ALL, diff --git a/lib/services/local/userpreferences_test.go b/lib/services/local/userpreferences_test.go index 8d2323820db..f156f127dc1 100644 --- a/lib/services/local/userpreferences_test.go +++ b/lib/services/local/userpreferences_test.go @@ -20,7 +20,6 @@ package local_test import ( "context" - "encoding/json" "testing" "github.com/google/go-cmp/cmp" @@ -30,7 +29,6 @@ import ( "google.golang.org/protobuf/testing/protocmp" userpreferencesv1 "github.com/gravitational/teleport/api/gen/proto/go/userpreferences/v1" - "github.com/gravitational/teleport/lib/backend" "github.com/gravitational/teleport/lib/backend/memory" "github.com/gravitational/teleport/lib/services/local" ) @@ -101,7 +99,6 @@ func TestUserPreferencesCRUD(t *testing.T) { }, }, expected: &userpreferencesv1.UserPreferences{ - Assist: defaultPref.Assist, Onboard: defaultPref.Onboard, Theme: userpreferencesv1.Theme_THEME_DARK, UnifiedResourcePreferences: defaultPref.UnifiedResourcePreferences, @@ -118,7 +115,6 @@ func TestUserPreferencesCRUD(t *testing.T) { }, }, expected: &userpreferencesv1.UserPreferences{ - Assist: defaultPref.Assist, Onboard: defaultPref.Onboard, Theme: defaultPref.Theme, UnifiedResourcePreferences: &userpreferencesv1.UnifiedResourcePreferences{ @@ -140,7 +136,6 @@ func TestUserPreferencesCRUD(t *testing.T) { }, }, expected: &userpreferencesv1.UserPreferences{ - Assist: defaultPref.Assist, Onboard: defaultPref.Onboard, Theme: defaultPref.Theme, UnifiedResourcePreferences: &userpreferencesv1.UnifiedResourcePreferences{ @@ -152,50 +147,6 @@ func TestUserPreferencesCRUD(t *testing.T) { ClusterPreferences: defaultPref.ClusterPreferences, }, }, - { - name: "update the assist preferred logins only", - req: &userpreferencesv1.UpsertUserPreferencesRequest{ - Preferences: &userpreferencesv1.UserPreferences{ - Assist: &userpreferencesv1.AssistUserPreferences{ - PreferredLogins: []string{"foo", "bar"}, - }, - Onboard: &userpreferencesv1.OnboardUserPreferences{ - PreferredResources: []userpreferencesv1.Resource{}, - MarketingParams: &userpreferencesv1.MarketingParams{}, - }, - }, - }, - expected: &userpreferencesv1.UserPreferences{ - Theme: defaultPref.Theme, - UnifiedResourcePreferences: defaultPref.UnifiedResourcePreferences, - Onboard: defaultPref.Onboard, - Assist: &userpreferencesv1.AssistUserPreferences{ - PreferredLogins: []string{"foo", "bar"}, - ViewMode: defaultPref.Assist.ViewMode, - }, - ClusterPreferences: defaultPref.ClusterPreferences, - }, - }, - { - name: "update the assist view mode only", - req: &userpreferencesv1.UpsertUserPreferencesRequest{ - Preferences: &userpreferencesv1.UserPreferences{ - Assist: &userpreferencesv1.AssistUserPreferences{ - ViewMode: userpreferencesv1.AssistViewMode_ASSIST_VIEW_MODE_POPUP_EXPANDED_SIDEBAR_VISIBLE, - }, - }, - }, - expected: &userpreferencesv1.UserPreferences{ - Theme: defaultPref.Theme, - UnifiedResourcePreferences: defaultPref.UnifiedResourcePreferences, - Onboard: defaultPref.Onboard, - Assist: &userpreferencesv1.AssistUserPreferences{ - PreferredLogins: defaultPref.Assist.PreferredLogins, - ViewMode: userpreferencesv1.AssistViewMode_ASSIST_VIEW_MODE_POPUP_EXPANDED_SIDEBAR_VISIBLE, - }, - ClusterPreferences: defaultPref.ClusterPreferences, - }, - }, { name: "update the onboard preference only", req: &userpreferencesv1.UpsertUserPreferencesRequest{ @@ -212,7 +163,6 @@ func TestUserPreferencesCRUD(t *testing.T) { }, }, expected: &userpreferencesv1.UserPreferences{ - Assist: defaultPref.Assist, Theme: defaultPref.Theme, UnifiedResourcePreferences: defaultPref.UnifiedResourcePreferences, Onboard: &userpreferencesv1.OnboardUserPreferences{ @@ -239,7 +189,6 @@ func TestUserPreferencesCRUD(t *testing.T) { }, }, expected: &userpreferencesv1.UserPreferences{ - Assist: defaultPref.Assist, Theme: defaultPref.Theme, UnifiedResourcePreferences: defaultPref.UnifiedResourcePreferences, Onboard: defaultPref.Onboard, @@ -261,10 +210,6 @@ func TestUserPreferencesCRUD(t *testing.T) { LabelsViewMode: userpreferencesv1.LabelsViewMode_LABELS_VIEW_MODE_COLLAPSED, AvailableResourceMode: userpreferencesv1.AvailableResourceMode_AVAILABLE_RESOURCE_MODE_NONE, }, - Assist: &userpreferencesv1.AssistUserPreferences{ - PreferredLogins: []string{"baz"}, - ViewMode: userpreferencesv1.AssistViewMode_ASSIST_VIEW_MODE_POPUP, - }, Onboard: &userpreferencesv1.OnboardUserPreferences{ PreferredResources: []userpreferencesv1.Resource{userpreferencesv1.Resource_RESOURCE_KUBERNETES}, MarketingParams: &userpreferencesv1.MarketingParams{ @@ -289,10 +234,6 @@ func TestUserPreferencesCRUD(t *testing.T) { LabelsViewMode: userpreferencesv1.LabelsViewMode_LABELS_VIEW_MODE_COLLAPSED, AvailableResourceMode: userpreferencesv1.AvailableResourceMode_AVAILABLE_RESOURCE_MODE_NONE, }, - Assist: &userpreferencesv1.AssistUserPreferences{ - PreferredLogins: []string{"baz"}, - ViewMode: userpreferencesv1.AssistViewMode_ASSIST_VIEW_MODE_POPUP, - }, Onboard: &userpreferencesv1.OnboardUserPreferences{ PreferredResources: []userpreferencesv1.Resource{userpreferencesv1.Resource_RESOURCE_KUBERNETES}, MarketingParams: &userpreferencesv1.MarketingParams{ @@ -335,37 +276,3 @@ func TestUserPreferencesCRUD(t *testing.T) { }) } } - -func TestLayoutUpdate(t *testing.T) { - t.Parallel() - - ctx := context.Background() - identity := newUserPreferencesService(t) - - outdatedPrefs := &userpreferencesv1.UserPreferences{ - Assist: &userpreferencesv1.AssistUserPreferences{ - PreferredLogins: []string{"foo", "bar"}, - }, - } - val, err := json.Marshal(outdatedPrefs) - require.NoError(t, err) - - // Insert the outdated preferences directly into the backend - // to simulate a previous version of the preferences. - _, err = identity.Put(ctx, backend.Item{ - Key: backend.Key("user_preferences", "test"), - Value: val, - }) - require.NoError(t, err) - - // Get the preferences and ensure that the layout is updated. - prefs, err := identity.GetUserPreferences(ctx, "test") - require.NoError(t, err) - // The layout should be updated to the latest version (values should not be nil). - require.NotNil(t, prefs.Onboard) - // Non-existing values should be set to the default value. - require.Equal(t, userpreferencesv1.AssistViewMode_ASSIST_VIEW_MODE_DOCKED, prefs.Assist.ViewMode) - require.Equal(t, userpreferencesv1.Theme_THEME_UNSPECIFIED, prefs.Theme) - // Existing values should be preserved. - require.Equal(t, []string{"foo", "bar"}, prefs.Assist.PreferredLogins) -} diff --git a/lib/services/local/users.go b/lib/services/local/users.go index 3d3fca348fd..8bf7a74b62b 100644 --- a/lib/services/local/users.go +++ b/lib/services/local/users.go @@ -1888,27 +1888,25 @@ func keyAttestationDataFingerprint(pubDER []byte) string { } const ( - webPrefix = "web" - usersPrefix = "users" - sessionsPrefix = "sessions" - attemptsPrefix = "attempts" - pwdPrefix = "pwd" - connectorsPrefix = "connectors" - oidcPrefix = "oidc" - samlPrefix = "saml" - githubPrefix = "github" - requestsPrefix = "requests" - requestsTracePrefix = "requestsTrace" - usedTOTPPrefix = "used_totp" - usedTOTPTTL = 30 * time.Second - mfaDevicePrefix = "mfa" - webauthnPrefix = "webauthn" - webauthnGlobalSessionData = "sessionData" - webauthnLocalAuthPrefix = "webauthnlocalauth" - webauthnSessionData = "webauthnsessiondata" - recoveryCodesPrefix = "recoverycodes" - attestationsPrefix = "key_attestations" - assistantMessagePrefix = "assistant_messages" - assistantConversationPrefix = "assistant_conversations" - userPreferencesPrefix = "user_preferences" + webPrefix = "web" + usersPrefix = "users" + sessionsPrefix = "sessions" + attemptsPrefix = "attempts" + pwdPrefix = "pwd" + connectorsPrefix = "connectors" + oidcPrefix = "oidc" + samlPrefix = "saml" + githubPrefix = "github" + requestsPrefix = "requests" + requestsTracePrefix = "requestsTrace" + usedTOTPPrefix = "used_totp" + usedTOTPTTL = 30 * time.Second + mfaDevicePrefix = "mfa" + webauthnPrefix = "webauthn" + webauthnGlobalSessionData = "sessionData" + webauthnLocalAuthPrefix = "webauthnlocalauth" + webauthnSessionData = "webauthnsessiondata" + recoveryCodesPrefix = "recoverycodes" + attestationsPrefix = "key_attestations" + userPreferencesPrefix = "user_preferences" ) diff --git a/lib/services/presets.go b/lib/services/presets.go index 9436728161d..1ca83f0d637 100644 --- a/lib/services/presets.go +++ b/lib/services/presets.go @@ -162,7 +162,6 @@ func NewPresetEditorRole() types.Role { types.NewRule(types.KindPlugin, RW()), types.NewRule(types.KindOktaImportRule, RW()), types.NewRule(types.KindOktaAssignment, RW()), - types.NewRule(types.KindAssistant, append(RW(), types.VerbUse)), types.NewRule(types.KindLock, RW()), types.NewRule(types.KindIntegration, append(RW(), types.VerbUse)), types.NewRule(types.KindBilling, RW()), @@ -234,7 +233,6 @@ func NewPresetAccessRole() types.Role { Where: "contains(session.participants, user.metadata.name)", }, types.NewRule(types.KindInstance, RO()), - types.NewRule(types.KindAssistant, append(RW(), types.VerbUse)), types.NewRule(types.KindClusterMaintenanceConfig, RO()), }, }, diff --git a/lib/services/useracl.go b/lib/services/useracl.go index 7443af2520e..8e9c0415685 100644 --- a/lib/services/useracl.go +++ b/lib/services/useracl.go @@ -88,8 +88,6 @@ type UserACL struct { DeviceTrust ResourceAccess `json:"deviceTrust"` // Locks defines access to locking resources. Locks ResourceAccess `json:"lock"` - // Assist defines access to assist feature. - Assist ResourceAccess `json:"assist"` // SAMLIdpServiceProvider defines access to `saml_idp_service_provider` objects. SAMLIdpServiceProvider ResourceAccess `json:"samlIdpServiceProvider"` // AccessList defines access to access list management. @@ -162,11 +160,6 @@ func NewUserACL(user types.User, userRoles RoleSet, features proto.Features, des activeSessionAccess.Read = true } - var assistAccess ResourceAccess - if features.Assist { - assistAccess = newAccess(userRoles, ctx, types.KindAssistant) - } - // The billing dashboards are available in cloud clusters or for // self-hosted dashboards for usage-based subscriptions. var billingAccess ResourceAccess @@ -237,7 +230,6 @@ func NewUserACL(user types.User, userRoles RoleSet, features proto.Features, des DiscoveryConfig: discoveryConfigsAccess, DeviceTrust: deviceTrust, Locks: lockAccess, - Assist: assistAccess, SAMLIdpServiceProvider: samlIdpServiceProviderAccess, AccessList: accessListAccess, AuditQuery: auditQuery, diff --git a/lib/teleterm/services/userpreferences/userpreferences_test.go b/lib/teleterm/services/userpreferences/userpreferences_test.go index 230da4a3c3f..c6583806a73 100644 --- a/lib/teleterm/services/userpreferences/userpreferences_test.go +++ b/lib/teleterm/services/userpreferences/userpreferences_test.go @@ -27,7 +27,6 @@ import ( ) var rootPreferencesMock = &userpreferencesv1.UserPreferences{ - Assist: nil, Onboard: nil, Theme: userpreferencesv1.Theme_THEME_LIGHT, ClusterPreferences: &userpreferencesv1.ClusterUserPreferences{ @@ -43,7 +42,6 @@ var rootPreferencesMock = &userpreferencesv1.UserPreferences{ } var leafPreferencesMock = &userpreferencesv1.UserPreferences{ - Assist: nil, Onboard: nil, ClusterPreferences: &userpreferencesv1.ClusterUserPreferences{ PinnedResources: &userpreferencesv1.PinnedResourcesUserPreferences{ diff --git a/lib/web/apiserver.go b/lib/web/apiserver.go index 45fdecac8c3..ef0c62902cf 100644 --- a/lib/web/apiserver.go +++ b/lib/web/apiserver.go @@ -49,14 +49,12 @@ import ( "github.com/gravitational/trace" "github.com/jonboulle/clockwork" "github.com/julienschmidt/httprouter" - "github.com/sashabaranov/go-openai" "github.com/sirupsen/logrus" "go.opentelemetry.io/otel/exporters/otlp/otlptrace" oteltrace "go.opentelemetry.io/otel/trace" tracepb "go.opentelemetry.io/proto/otlp/trace/v1" "golang.org/x/crypto/ssh" "golang.org/x/mod/semver" - "golang.org/x/time/rate" "google.golang.org/protobuf/encoding/protojson" "google.golang.org/protobuf/types/known/timestamppb" @@ -115,14 +113,6 @@ const ( // callback URL in tsh login. SSOLoginFailureInvalidRedirect = "Failed to login due to a disallowed callback URL. Please check Teleport's log for more details." - // assistantTokensPerHour defines how many assistant rate limiter tokens are replenished every hour. - assistantTokensPerHour = 140 - // assistantLimiterRate is the rate (in tokens per second) - // at which tokens for the assistant rate limiter are replenished - assistantLimiterRate = rate.Limit(assistantTokensPerHour / float64(time.Hour/time.Second)) - // assistantLimiterCapacity is the total capacity of the token bucket for the assistant rate limiter. - // The bucket starts full, prefilled for a week. - assistantLimiterCapacity = assistantTokensPerHour * 24 * 7 // webUIFlowLabelKey is a label that may be added to resources // created via the web UI, indicating which flow the resource was created on. // This label is used for enhancing UX in the web app, by showing icons related, @@ -155,13 +145,7 @@ type Handler struct { clock clockwork.Clock limiter *limiter.RateLimiter highLimiter *limiter.RateLimiter - // assistantLimiter limits the amount of tokens that can be consumed - // by OpenAI API calls when using a shared key. - // golang.org/x/time/rate is used, as the oxy ratelimiter - // is quite tightly tied to individual http.Requests, - // and instead we want to consume arbitrary amounts of tokens. - assistantLimiter *rate.Limiter - healthCheckAppServer healthCheckAppServerFunc + healthCheckAppServer healthCheckAppServerFunc // sshPort specifies the SSH proxy port extracted // from configuration sshPort string @@ -312,9 +296,6 @@ type Config struct { // UI provides config options for the web UI UI webclient.UIConfig - // OpenAIConfig provides config options for the OpenAI integration. - OpenAIConfig *openai.ClientConfig - // NodeWatcher is a services.NodeWatcher used by Assist to lookup nodes from // the proxy's cache and get nodes in real time. NodeWatcher *services.NodeWatcher @@ -409,15 +390,6 @@ func NewHandler(cfg Config, opts ...HandlerOption) (*APIHandler, error) { wsIODeadline: wsIODeadline, } - // Check for self-hosted vs Cloud. - // TODO(justinas): this needs to be modified when we allow user-supplied API keys in Cloud - if cfg.ClusterFeatures.GetCloud() { - h.assistantLimiter = rate.NewLimiter(assistantLimiterRate, assistantLimiterCapacity) - } else { - // Set up a limiter with "infinite limit", the "burst" parameter is ignored - h.assistantLimiter = rate.NewLimiter(rate.Inf, 0) - } - if automaticUpgrades(cfg.ClusterFeatures) && h.cfg.AutomaticUpgradesChannels == nil { h.cfg.AutomaticUpgradesChannels = automaticupgrades.Channels{} } @@ -960,9 +932,6 @@ func (h *Handler) bindDefaultEndpoints() { h.GET("/webapi/sites/:site/user-groups", h.WithClusterAuth(h.getUserGroups)) - // WebSocket endpoint for the chat conversation, websocket auth - h.GET("/webapi/sites/:site/assistant/ws", h.WithClusterAuthWebSocket(h.assistant)) - // Fetches the user's preferences h.GET("/webapi/user/preferences", h.WithAuth(h.getUserPreferences)) @@ -1436,17 +1405,6 @@ func (h *Handler) ping(w http.ResponseWriter, r *http.Request, p httprouter.Para return nil, trace.Wrap(err) } - // TODO(jakule): This part should be removed once the plugin support is added to OSS. - if proxyConfig.AssistEnabled { - enabled, err := h.cfg.ProxyClient.IsAssistEnabled(r.Context()) - if err != nil { - return webclient.AuthenticationSettings{}, trace.Wrap(err) - } - - // disable if auth doesn't support assist - proxyConfig.AssistEnabled = enabled.Enabled - } - pr, err := h.cfg.ProxyClient.Ping(r.Context()) if err != nil { return nil, trace.Wrap(err) @@ -1655,7 +1613,6 @@ func (h *Handler) getWebConfig(w http.ResponseWriter, r *http.Request, p httprou // get tunnel address to display on cloud instances tunnelPublicAddr := "" - assistEnabled := false // TODO(jakule) remove when plugins are implemented proxyConfig, err := h.cfg.ProxySettings.GetProxySettings(r.Context()) if err != nil { h.log.WithError(err).Warn("Cannot retrieve ProxySettings, tunnel address won't be set in Web UI.") @@ -1663,16 +1620,6 @@ func (h *Handler) getWebConfig(w http.ResponseWriter, r *http.Request, p httprou if clusterFeatures.GetCloud() { tunnelPublicAddr = proxyConfig.SSH.TunnelPublicAddr } - // TODO(jakule): This part should be removed once the plugin support is added to OSS. - if proxyConfig.AssistEnabled { - enabled, err := h.cfg.ProxyClient.IsAssistEnabled(r.Context()) - if err != nil { - return nil, trace.Wrap(err) - } - - // disable if auth doesn't support assist - assistEnabled = enabled.Enabled - } } // disable joining sessions if proxy session recording is enabled @@ -1710,7 +1657,6 @@ func (h *Handler) getWebConfig(w http.ResponseWriter, r *http.Request, p httprou IsUsageBasedBilling: clusterFeatures.GetIsUsageBased(), AutomaticUpgrades: automaticUpgradesEnabled, AutomaticUpgradesTargetVersion: automaticUpgradesTargetVersion, - AssistEnabled: assistEnabled, HideInaccessibleFeatures: clusterFeatures.GetFeatureHiding(), CustomTheme: clusterFeatures.GetCustomTheme(), IsIGSEnabled: clusterFeatures.GetIdentityGovernance(), diff --git a/lib/web/apiserver_test.go b/lib/web/apiserver_test.go index f2edf1232e6..bd177bbca48 100644 --- a/lib/web/apiserver_test.go +++ b/lib/web/apiserver_test.go @@ -58,7 +58,6 @@ import ( "github.com/julienschmidt/httprouter" "github.com/mailgun/timetools" "github.com/pquerna/otp/totp" - "github.com/sashabaranov/go-openai" "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -197,9 +196,6 @@ type webSuiteConfig struct { // services. HealthCheckAppServer healthCheckAppServerFunc - // OpenAIConfig is a custom OpenAI config for the test. - OpenAIConfig *openai.ClientConfig - // ClusterFeatures allows overriding default auth server features ClusterFeatures *authproto.Features @@ -493,7 +489,6 @@ func newWebSuiteWithConfig(t *testing.T, cfg webSuiteConfig) *WebSuite { Router: router, HealthCheckAppServer: cfg.HealthCheckAppServer, UI: cfg.uiConfig, - OpenAIConfig: cfg.OpenAIConfig, PresenceChecker: cfg.presenceChecker, GetProxyIdentity: func() (*state.Identity, error) { return proxyIdentity, nil @@ -4578,7 +4573,6 @@ func TestGetWebConfig(t *testing.T) { CanJoinSessions: true, ProxyClusterName: env.server.ClusterName(), IsCloud: false, - AssistEnabled: false, AutomaticUpgrades: false, JoinActiveSessions: true, Edition: modules.BuildOSS, // testBuildType is empty @@ -4608,13 +4602,6 @@ func TestGetWebConfig(t *testing.T) { }, }) - mockProxySetting := &mockProxySettings{ - mockedGetProxySettings: func(ctx context.Context) (*webclient.ProxySettings, error) { - return &webclient.ProxySettings{AssistEnabled: true}, nil - }, - } - env.proxies[0].handler.handler.cfg.ProxySettings = mockProxySetting - require.NoError(t, err) // This version is too high and MUST NOT be used testVersion := "v99.0.1" @@ -4630,7 +4617,6 @@ func TestGetWebConfig(t *testing.T) { expectedCfg.IsUsageBasedBilling = true expectedCfg.AutomaticUpgrades = true expectedCfg.AutomaticUpgradesTargetVersion = "v" + teleport.Version - expectedCfg.AssistEnabled = false expectedCfg.JoinActiveSessions = false expectedCfg.Edition = "" // testBuildType is empty diff --git a/lib/web/assistant.go b/lib/web/assistant.go deleted file mode 100644 index 23fd3852d5d..00000000000 --- a/lib/web/assistant.go +++ /dev/null @@ -1,371 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package web - -import ( - "context" - "errors" - "io" - "net/http" - "time" - - "github.com/gorilla/websocket" - "github.com/gravitational/trace" - "github.com/julienschmidt/httprouter" - - "github.com/gravitational/teleport/api/client/proto" - assistpb "github.com/gravitational/teleport/api/gen/proto/go/assist/v1" - usageeventsv1 "github.com/gravitational/teleport/api/gen/proto/go/usageevents/v1" - "github.com/gravitational/teleport/lib/ai/model/tools" - "github.com/gravitational/teleport/lib/ai/tokens" - "github.com/gravitational/teleport/lib/assist" - "github.com/gravitational/teleport/lib/auth/authclient" - "github.com/gravitational/teleport/lib/reversetunnelclient" -) - -const ( - // actionSSHGenerateCommand is a name of the action for generating SSH commands. - actionSSHGenerateCommand = "ssh-cmdgen" - // actionSSHExplainCommand is a name of the action for explaining terminal output in SSH session. - actionSSHExplainCommand = "ssh-explain" - // actionGenerateAuditQuery is the name of the action for generating audit queries. - actionGenerateAuditQuery = "audit-query" - // We cannot know how many tokens we will consume in advance. - // Try to consume a small number of tokens first. - lookaheadTokens = 100 -) - -// assistantMessage is an assistant message that is sent to the client. -type assistantMessage struct { - // Type is a type of the message. - Type assist.MessageType `json:"type"` - // CreatedTime is a time when the message was created in RFC3339 format. - CreatedTime string `json:"created_time"` - // Payload is a message payload in JSON format. - Payload string `json:"payload"` -} - -// assistant is a handler for GET /webapi/sites/:site/assistant. -// This handler covers the main chat conversation as well as the -// SSH competition (SSH command generation and output explanation). -func (h *Handler) assistant(w http.ResponseWriter, r *http.Request, _ httprouter.Params, - sctx *SessionContext, site reversetunnelclient.RemoteSite, ws *websocket.Conn, -) (any, error) { - if err := runAssistant(h, w, r, sctx, site, ws); err != nil { - h.log.Warn(trace.DebugReport(err)) - return nil, trace.Wrap(err) - } - - return nil, nil -} - -// reserveTokens preemptively reserves tokens in the ratelimiter. -func (h *Handler) reserveTokens(usedTokens *tokens.TokenCount) (int, int) { - promptTokens, completionTokens := usedTokens.CountAll() - - // Once we know how many tokens were consumed for prompt+completion, - // consume the remaining tokens from the rate limiter bucket. - extraTokens := promptTokens + completionTokens - lookaheadTokens - if extraTokens < 0 { - extraTokens = 0 - } - h.assistantLimiter.ReserveN(time.Now(), extraTokens) - return promptTokens, completionTokens -} - -// reportTokenUsage sends a token usage event for an action. -func (h *Handler) reportActionTokenUsage(authClient authclient.ClientI, usedTokens *tokens.TokenCount, action string) { - // Create a new context to not be bounded by the request timeout. - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - - promptTokens, completionTokens := h.reserveTokens(usedTokens) - usageEventReq := &proto.SubmitUsageEventRequest{ - Event: &usageeventsv1.UsageEventOneOf{ - Event: &usageeventsv1.UsageEventOneOf_AssistAction{ - AssistAction: &usageeventsv1.AssistAction{ - Action: action, - TotalTokens: int64(promptTokens + completionTokens), - PromptTokens: int64(promptTokens), - CompletionTokens: int64(completionTokens), - }, - }, - }, - } - - if err := authClient.SubmitUsageEvent(ctx, usageEventReq); err != nil { - h.log.WithError(err).Warn("Failed to emit usage event") - } -} - -func checkAssistEnabled(a authclient.ClientI, ctx context.Context) error { - enabled, err := a.IsAssistEnabled(ctx) - if err != nil { - return trace.Wrap(err) - } - if !enabled.Enabled { - return trace.AccessDenied("Assist is not enabled") - } - - return nil -} - -// runAssistant upgrades the HTTP connection to a websocket and starts a chat loop. -func runAssistant(h *Handler, w http.ResponseWriter, r *http.Request, - sctx *SessionContext, site reversetunnelclient.RemoteSite, ws *websocket.Conn, -) (err error) { - q := r.URL.Query() - conversationID := q.Get("conversation_id") - actionParam := r.URL.Query().Get("action") - if conversationID == "" && actionParam == "" { - return trace.BadParameter("conversation ID or action is required") - } - - authClient, err := sctx.GetClient() - if err != nil { - return trace.Wrap(err) - } - - if err := checkAssistEnabled(authClient, r.Context()); err != nil { - return trace.Wrap(err) - } - - ctx, err := h.cfg.SessionControl.AcquireSessionContext(r.Context(), sctx, sctx.GetUser(), h.cfg.ProxyWebAddr.Addr, r.RemoteAddr) - if err != nil { - return trace.Wrap(err) - } - - authAccessPoint, err := site.CachingAccessPoint() - if err != nil { - h.log.WithError(err).Debug("Unable to get auth access point.") - return trace.Wrap(err) - } - - netConfig, err := authAccessPoint.GetClusterNetworkingConfig(ctx) - if err != nil { - h.log.WithError(err).Debug("Unable to fetch cluster networking config.") - return trace.Wrap(err) - } - - // Note: This time should be longer than OpenAI response time. - keepAliveInterval := netConfig.GetKeepAliveInterval() - err = ws.SetReadDeadline(deadlineForInterval(keepAliveInterval)) - if err != nil { - h.log.WithError(err).Error("Error setting websocket readline") - return nil - } - defer func() { - closureReason := websocket.CloseNormalClosure - closureMsg := "" - if err != nil { - h.log.WithError(err).Error("Error in the Assistant loop") - _ = ws.WriteJSON(&assistantMessage{ - Type: assist.MessageKindError, - Payload: "An error has occurred. Please try again later.", - CreatedTime: h.clock.Now().UTC().Format(time.RFC3339), - }) - // Set server error code and message: https://datatracker.ietf.org/doc/html/rfc6455#section-7.4.1 - closureReason = websocket.CloseInternalServerErr - closureMsg = err.Error() - } - // Send the close message to the client and close the connection - if err := ws.WriteControl(websocket.CloseMessage, - websocket.FormatCloseMessage(closureReason, closureMsg), - time.Now().Add(time.Second), - ); err != nil { - h.log.Warnf("Failed to write close message: %v", err) - } - if err := ws.Close(); err != nil { - h.log.Warnf("Failed to close websocket: %v", err) - } - }() - - // Update the read deadline upon receiving a pong message. - ws.SetPongHandler(func(_ string) error { - return trace.Wrap(ws.SetReadDeadline(deadlineForInterval(keepAliveInterval))) - }) - - ws.SetCloseHandler(func(code int, text string) error { - h.log.Debugf("closing assistant websocket: %v %v", code, text) - return nil - }) - - go startWSPingLoop(ctx, ws, keepAliveInterval, h.log, nil) - - assistClient, err := assist.NewClient(ctx, h.cfg.ProxyClient, - h.cfg.ProxySettings, h.cfg.OpenAIConfig) - if err != nil { - return trace.Wrap(err) - } - - switch r.URL.Query().Get("action") { - case actionSSHGenerateCommand: - err = h.assistGenSSHCommandLoop(ctx, assistClient, authClient, ws, sctx.GetUser()) - case actionSSHExplainCommand: - err = h.assistSSHExplainOutputLoop(ctx, assistClient, authClient, ws) - case actionGenerateAuditQuery: - err = h.assistGenAuditQueryLoop(ctx, assistClient, authClient, ws, sctx.GetUser()) - default: - err = trace.Errorf("Teleport Assist Chat has been remove in v16") - } - - return trace.Wrap(err) -} - -// assistGenAuditQueryLoop reads the user's input and generates an audit query. -func (h *Handler) assistGenAuditQueryLoop(ctx context.Context, assistClient *assist.Assist, authClient authclient.ClientI, ws *websocket.Conn, username string) error { - for { - _, payload, err := ws.ReadMessage() - if err != nil { - if wsIsClosed(err) { - break - } - return trace.Wrap(err) - } - - onMessage := func(kind assist.MessageType, payload []byte, createdTime time.Time) error { - return onMessageFn(ws, kind, payload, createdTime) - } - - toolCtx := &tools.ToolContext{User: username} - - if err := h.preliminaryRateLimitGuard(onMessage); err != nil { - return trace.Wrap(err) - } - - tokenCount, err := assistClient.RunTool(ctx, onMessage, tools.AuditQueryGenerationToolName, string(payload), toolCtx) - if err != nil { - return trace.Wrap(err) - } - - go h.reportActionTokenUsage(authClient, tokenCount, tools.AuditQueryGenerationToolName) - } - return nil -} - -// assistSSHExplainOutputLoop reads the user's input and generates a command summary. -func (h *Handler) assistSSHExplainOutputLoop(ctx context.Context, assistClient *assist.Assist, authClient authclient.ClientI, ws *websocket.Conn) error { - _, payload, err := ws.ReadMessage() - if err != nil { - if wsIsClosed(err) { - return nil - } - return trace.Wrap(err) - } - - modelMessages := []*assistpb.AssistantMessage{ - { - Type: string(assist.MessageKindUserMessage), - Payload: string(payload), - }, - } - - onMessage := func(kind assist.MessageType, payload []byte, createdTime time.Time) error { - return onMessageFn(ws, kind, payload, createdTime) - } - - if err := h.preliminaryRateLimitGuard(onMessage); err != nil { - return trace.Wrap(err) - } - - summary, tokenCount, err := assistClient.GenerateCommandSummary(ctx, modelMessages, map[string][]byte{}) - if err != nil { - return trace.Wrap(err) - } - - if err := onMessageFn(ws, assist.MessageKindAssistantMessage, []byte(summary), h.clock.Now().UTC()); err != nil { - return trace.Wrap(err) - } - - go h.reportActionTokenUsage(authClient, tokenCount, "SSH Explain") - return nil -} - -// assistSSHCommandLoop reads the user's input and generates a Linux command. -func (h *Handler) assistGenSSHCommandLoop(ctx context.Context, assistClient *assist.Assist, authClient authclient.ClientI, ws *websocket.Conn, username string) error { - chat, err := assistClient.NewLightweightChat(username) - if err != nil { - return trace.Wrap(err) - } - - onMessage := func(kind assist.MessageType, payload []byte, createdTime time.Time) error { - return trace.Wrap(onMessageFn(ws, kind, payload, createdTime)) - } - - for { - _, payload, err := ws.ReadMessage() - if err != nil { - if wsIsClosed(err) { - break - } - return trace.Wrap(err) - } - - if err := h.preliminaryRateLimitGuard(onMessage); err != nil { - return trace.Wrap(err) - } - - tokenCount, err := chat.ProcessComplete(ctx, func(kind assist.MessageType, payload []byte, createdTime time.Time) error { - return onMessageFn(ws, kind, payload, createdTime) - }, string(payload)) - if err != nil { - return trace.Wrap(err) - } - - tool := tools.CommandExecutionTool{} - go h.reportActionTokenUsage(authClient, tokenCount, tool.Name()) - } - return nil -} - -// preliminaryRateLimitGuard checks that some small number of tokens is still available and the ratelimit is not exceeded. -// This is done because the changed quantity within the limiter is not known until after a request is processed. -func (h *Handler) preliminaryRateLimitGuard(onMessageFn func(kind assist.MessageType, payload []byte, createdTime time.Time) error) error { - const errorMsg = "You have reached the rate limit. Please try again later." - - if !h.assistantLimiter.AllowN(time.Now(), lookaheadTokens) { - err := onMessageFn(assist.MessageKindError, []byte(errorMsg), h.clock.Now().UTC()) - if err != nil { - return trace.Wrap(err) - } - - return trace.LimitExceeded(errorMsg) - } - - return nil -} - -// wsIsClosed returns true if the error is caused by a closed websocket. -func wsIsClosed(err error) bool { - return errors.Is(err, io.EOF) || websocket.IsCloseError(err, websocket.CloseAbnormalClosure, - websocket.CloseGoingAway, websocket.CloseNormalClosure) -} - -// onMessageFn is a helper function used to send an assist message to the frontend. -// It deals with serializing the kind and payload into a wire and sending it over with the correct -// websocket frame type. -func onMessageFn(ws *websocket.Conn, kind assist.MessageType, payload []byte, createdTime time.Time) error { - msg := &assistantMessage{ - Type: kind, - Payload: string(payload), - CreatedTime: createdTime.Format(time.RFC3339), - } - - return trace.Wrap(ws.WriteJSON(msg)) -} diff --git a/lib/web/assistant_test.go b/lib/web/assistant_test.go deleted file mode 100644 index bcc823e0e95..00000000000 --- a/lib/web/assistant_test.go +++ /dev/null @@ -1,203 +0,0 @@ -/* - * Teleport - * Copyright (C) 2023 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package web - -import ( - "crypto/tls" - "encoding/json" - "fmt" - "net/http" - "net/http/httptest" - "net/url" - "testing" - - "github.com/gorilla/websocket" - "github.com/gravitational/trace" - "github.com/sashabaranov/go-openai" - "github.com/stretchr/testify/require" - - "github.com/gravitational/teleport/api/types" - aitest "github.com/gravitational/teleport/lib/ai/testutils" - "github.com/gravitational/teleport/lib/assist" - "github.com/gravitational/teleport/lib/client" - "github.com/gravitational/teleport/lib/services" -) - -func Test_SSHCommandGeneration(t *testing.T) { - t.Parallel() - - assertGenCommand := func(ws *websocket.Conn) { - _, payload, err := ws.ReadMessage() - require.NoError(t, err) - - var msg assistantMessage - err = json.Unmarshal(payload, &msg) - require.NoError(t, err) - - // Expect "hello" message - require.Equal(t, assist.MessageKindProgressUpdate, msg.Type) - require.Contains(t, msg.Payload, "openssl req -x509 -newkey rsa:4096 -keyout key.pem -out cert.pem -days 365 -nodes -subj") - } - - responses := []string{generateCommandResponse()} - server := httptest.NewServer(aitest.GetTestHandlerFn(t, responses)) - t.Cleanup(server.Close) - - openaiCfg := openai.DefaultConfig("test-token") - openaiCfg.BaseURL = server.URL - s := newWebSuiteWithConfig(t, webSuiteConfig{OpenAIConfig: &openaiCfg}) - - assistRole := allowAssistAccess(t, s) - authPack := s.authPack(t, "foo", assistRole.GetName()) - - // Make WS client and start the conversation - ws, err := s.makeAssistant(t, authPack, "", "ssh-cmdgen") - require.NoError(t, err) - t.Cleanup(func() { - err := ws.Close() - require.NoError(t, err) - }) - - err = ws.WriteMessage(websocket.TextMessage, []byte(`{"input:" "My cert expired!!! What is x509?"}`)) - require.NoError(t, err) - - // verify responses - assertGenCommand(ws) -} - -func Test_SSHCommandExplain(t *testing.T) { - t.Parallel() - - assertResponse := func(ws *websocket.Conn) { - _, payload, err := ws.ReadMessage() - require.NoError(t, err) - - var msg assistantMessage - err = json.Unmarshal(payload, &msg) - require.NoError(t, err) - - // Expect "hello" message - require.Equal(t, assist.MessageKindAssistantMessage, msg.Type) - require.Contains(t, msg.Payload, "The application has failed to connect to the database. The database is not running.") - } - - responses := []string{commandSummaryResponse()} - server := httptest.NewServer(aitest.GetTestHandlerFn(t, responses)) - t.Cleanup(server.Close) - - openaiCfg := openai.DefaultConfig("test-token") - openaiCfg.BaseURL = server.URL - s := newWebSuiteWithConfig(t, webSuiteConfig{OpenAIConfig: &openaiCfg}) - - assistRole := allowAssistAccess(t, s) - authPack := s.authPack(t, "foo", assistRole.GetName()) - - // Make WS client and start the conversation - ws, err := s.makeAssistant(t, authPack, "", "ssh-explain") - require.NoError(t, err) - t.Cleanup(func() { - err := ws.Close() - require.NoError(t, err) - }) - - err = ws.WriteMessage(websocket.TextMessage, []byte(`{"input:" "listen tcp 0.0.0.0:5432: bind: address already in use"}`)) - require.NoError(t, err) - - // verify responses - assertResponse(ws) -} - -func allowAssistAccess(t *testing.T, s *WebSuite) types.Role { - assistRole, err := types.NewRole("assist-access", types.RoleSpecV6{ - Allow: types.RoleConditions{ - Rules: []types.Rule{ - types.NewRule(types.KindAssistant, services.RW()), - }, - }, - }) - require.NoError(t, err) - assistRole, err = s.server.Auth().UpsertRole(s.ctx, assistRole) - require.NoError(t, err) - - return assistRole -} - -// makeAssistant creates a new assistant websocket connection. -func (s *WebSuite) makeAssistant(_ *testing.T, pack *authPack, conversationID, action string) (*websocket.Conn, error) { - if action == "" && conversationID == "" { - return nil, trace.BadParameter("must specify either conversation_id or action") - } - - u := url.URL{ - Host: s.url().Host, - Scheme: client.WSS, - Path: fmt.Sprintf("/v1/webapi/sites/%s/assistant/ws", currentSiteShortcut), - } - - q := u.Query() - if conversationID != "" { - q.Set("conversation_id", conversationID) - } - - if action != "" { - q.Set("action", action) - } - - u.RawQuery = q.Encode() - - dialer := websocket.Dialer{} - dialer.TLSClientConfig = &tls.Config{ - InsecureSkipVerify: true, - } - - header := http.Header{} - header.Add("Origin", "http://localhost") - for _, cookie := range pack.cookies { - header.Add("Cookie", cookie.String()) - } - - ws, resp, err := dialer.Dial(u.String(), header) - if err != nil { - return nil, trace.Wrap(err) - } - - if err := makeAuthReqOverWS(ws, pack.session.Token); err != nil { - return nil, trace.Wrap(err) - } - - err = resp.Body.Close() - if err != nil { - return nil, trace.Wrap(err) - } - - return ws, nil -} - -func generateCommandResponse() string { - return "```" + `json - { - "action": "Command Generation", - "action_input": "{\"command\":\"openssl req -x509 -newkey rsa:4096 -keyout key.pem -out cert.pem -days 365 -nodes -subj\"}" - } - ` + "```" -} - -func commandSummaryResponse() string { - return "The application has failed to connect to the database. The database is not running." -} diff --git a/lib/web/userpreferences.go b/lib/web/userpreferences.go index a86acd4da11..f4e8517449b 100644 --- a/lib/web/userpreferences.go +++ b/lib/web/userpreferences.go @@ -135,10 +135,6 @@ func makePreferenceRequest(req UserPreferencesResponse) *userpreferencesv1.Upser LabelsViewMode: req.UnifiedResourcePreferences.LabelsViewMode, AvailableResourceMode: req.UnifiedResourcePreferences.AvailableResourceMode, }, - Assist: &userpreferencesv1.AssistUserPreferences{ - PreferredLogins: req.Assist.PreferredLogins, - ViewMode: req.Assist.ViewMode, - }, Onboard: &userpreferencesv1.OnboardUserPreferences{ PreferredResources: req.Onboard.PreferredResources, MarketingParams: &userpreferencesv1.MarketingParams{ @@ -184,7 +180,6 @@ func (h *Handler) updateUserPreferences(_ http.ResponseWriter, r *http.Request, // userPreferencesResponse creates a JSON response for the user preferences. func userPreferencesResponse(resp *userpreferencesv1.UserPreferences) *UserPreferencesResponse { jsonResp := &UserPreferencesResponse{ - Assist: assistUserPreferencesResponse(resp.Assist), Theme: resp.Theme, Onboard: onboardUserPreferencesResponse(resp.Onboard), ClusterPreferences: clusterPreferencesResponse(resp.ClusterPreferences), @@ -206,18 +201,6 @@ func clusterPreferencesResponse(prefs *userpreferencesv1.ClusterUserPreferences) return resp } -// assistUserPreferencesResponse creates a JSON response for the assist user preferences. -func assistUserPreferencesResponse(resp *userpreferencesv1.AssistUserPreferences) AssistUserPreferencesResponse { - jsonResp := AssistUserPreferencesResponse{ - PreferredLogins: make([]string, 0, len(resp.PreferredLogins)), - ViewMode: resp.ViewMode, - } - - jsonResp.PreferredLogins = append(jsonResp.PreferredLogins, resp.PreferredLogins...) - - return jsonResp -} - // unifiedResourcePreferencesResponse creates a JSON response for the assist user preferences. func unifiedResourcePreferencesResponse(resp *userpreferencesv1.UnifiedResourcePreferences) UnifiedResourcePreferencesResponse { return UnifiedResourcePreferencesResponse{ diff --git a/tool/tctl/common/resource_command_test.go b/tool/tctl/common/resource_command_test.go index a312b21fd4b..613be1abf0d 100644 --- a/tool/tctl/common/resource_command_test.go +++ b/tool/tctl/common/resource_command_test.go @@ -1741,7 +1741,6 @@ func testCreateClusterNetworkingConfig(t *testing.T, fc *config.FileConfig) { metadata: name: cluster-networking-config spec: - assist_command_execution_workers: 30 client_idle_timeout: 0s idle_timeout_message: "" keep_alive_count_max: 300