mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
Remove Assist (#42657)
* Remove assist feature * Fix some tests * Remove unused functions * Fix ut Remove more stuff
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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
@@ -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);
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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{}
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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:]
|
||||
}
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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\":[]}"
|
||||
}
|
||||
` + "```"
|
||||
}
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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()),
|
||||
},
|
||||
},
|
||||
|
||||
@@ -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
@@ -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(),
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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."
|
||||
}
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user