Remove Assist (#42657)

* Remove assist feature

* Fix some tests

* Remove unused functions

* Fix ut

Remove more stuff
This commit is contained in:
Jakub Nyckowski
2024-06-10 20:57:05 +00:00
committed by GitHub
parent 78d6a38783
commit bfaa3340e8
86 changed files with 143 additions and 11082 deletions
-17
View File
@@ -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.
+2 -3
View File
@@ -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
-81
View File
@@ -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)
-2
View File
@@ -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
-2
View File
@@ -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.
-4
View File
@@ -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 (
File diff suppressed because it is too large Load Diff
@@ -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",
}
@@ -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()
-201
View File
@@ -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);
}
@@ -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.
-4
View File
@@ -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"
-22
View File
@@ -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
+2 -5
View File
@@ -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
@@ -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<UserPreferences> {
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<UserPreferences> {
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<UserPreferences> {
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);
-1
View File
@@ -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
-2
View File
@@ -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=
@@ -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 {
-92
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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{}
}
-364
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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": "<FINAL RESPONSE>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())
}
-253
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-112
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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)
}
-115
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-152
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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)
}
-331
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-313
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-54
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-436
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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 = "<FINAL RESPONSE>"
)
// 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
}
}
-60
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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)
}
-79
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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"`
}
-48
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-83
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-131
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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())
}
-307
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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"`
}
-177
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-82
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-147
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-140
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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)
})
}
}
-84
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-77
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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)
}
-146
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-103
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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")
}
}
-131
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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:]
}
-31
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
)
-194
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-96
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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())
}
-623
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
}
-209
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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\":[]}"
}
` + "```"
}
-37
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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",
}
-36
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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)
}
-533
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
}
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-54
View File
@@ -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) {
-5
View File
@@ -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
-16
View File
@@ -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 {
+2 -26
View File
@@ -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
-14
View File
@@ -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
@@ -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{},
+11 -40
View File
@@ -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
}
-59
View File
@@ -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()
+1 -33
View File
@@ -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.
-4
View File
@@ -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,
}
-7
View File
@@ -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,
},
-111
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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)
})
}
}
-62
View File
@@ -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,
-4
View File
@@ -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
}
-10
View File
@@ -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
-4
View File
@@ -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.
-48
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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)
}
-40
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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)
}
-285
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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
}
-176
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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")
})
}
-117
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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())
}
-237
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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)
}
-4
View File
@@ -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,
@@ -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)
}
+21 -23
View File
@@ -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"
)
-2
View File
@@ -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()),
},
},
-8
View File
@@ -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,
@@ -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{
+1 -55
View File
@@ -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(),
-14
View File
@@ -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
-371
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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))
}
-203
View File
@@ -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 <http://www.gnu.org/licenses/>.
*/
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."
}
-17
View File
@@ -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{
@@ -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